Spaces:
Running on Zero
Running on Zero
File size: 7,299 Bytes
de2e2e5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
from gemmasight.config import MEDGEMMA_ID, FORCE_SIMULATION
class MedGemmaReportGenerator:
def __init__(self, force_simulation=FORCE_SIMULATION):
self.force_simulation = force_simulation
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if self.force_simulation:
print("GemmaSight MedGemma Generator: Initializing in Simulation/Fallback Mode...")
self.is_simulation = True
else:
try:
print("GemmaSight MedGemma Generator: Loading google/medgemma-1.5-4b-it from Hugging Face...")
self.tokenizer = AutoTokenizer.from_pretrained(MEDGEMMA_ID)
self.model = AutoModelForCausalLM.from_pretrained(
MEDGEMMA_ID,
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
device_map="auto" if torch.cuda.is_available() else None,
trust_remote_code=True
)
if not torch.cuda.is_available():
self.model.to(self.device)
self.model.eval()
print("GemmaSight MedGemma Generator: Model loaded successfully.")
self.is_simulation = False
except Exception as e:
print(f"GemmaSight MedGemma Generator Error: Failed to load: {e}")
print("GemmaSight MedGemma Generator: Falling back to Clinical Simulation Engine.")
self.is_simulation = True
def generate_report(self, probability: float, retrieved_cases: list) -> str:
"""
Generates a pathology report based on probability score and retrieved FAISS clinical cases.
"""
# Format similar historical cases for the prompt context
cases_text = ""
for i, c in enumerate(retrieved_cases):
cases_text += f"Case {i+1} (ID: {c['patient_id']}): {c['status']} with {c['similarity']:.1%} similarity. Feature findings: {c['visual_description']}\n"
# If simulation, run the medical rule-based template engine
if self.is_simulation:
return self._generate_simulated_report(probability, retrieved_cases)
# Live generation using MedGemma
prompt = f"""You are a clinical pathologist. Analyze the provided H&E tissue patch with the
overlay heatmap (highlighting regions most influential for the prediction).
Prediction: {probability:.2%} probability of MSI-High phenotype.
Similar historical cases:
{cases_text}
Generate a concise pathology report including:
1. Morphological observations from the heatmap-highlighted regions
2. Correlation with historical matches
3. Clinical interpretation of the MSI status prediction
4. Confidence level and limitations
Keep the report under 200 words, professional tone, evidence-based."""
try:
inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device)
with torch.no_grad():
outputs = self.model.generate(
**inputs,
max_new_tokens=250,
temperature=0.4,
top_p=0.9,
do_sample=True,
pad_token_id=self.tokenizer.eos_token_id
)
generated_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
# Remove prompt if it gets printed
if prompt in generated_text:
generated_text = generated_text.replace(prompt, "").strip()
return generated_text
except Exception as e:
print(f"MedGemma Inference Error: {e}. Falling back to simulated clinical writer.")
return self._generate_simulated_report(probability, retrieved_cases)
def _generate_simulated_report(self, probability: float, retrieved_cases: list) -> str:
"""
Generates an highly realistic, clinical-grade pathology report using rule-based templates.
"""
is_msi_high = probability >= 0.5
status_str = "MSI-High (Microsatellite Instability-High)" if is_msi_high else "MSS (Microsatellite Stable)"
# Primary morphological highlights based on predicted status
if is_msi_high:
morphology = (
"Heatmap highlights intense focal regions corresponding to abundant tumor-infiltrating lymphocytes (TILs) "
"and distinct Crohn's-like lymphoid aggregates. Glandular structures demonstrate significant mucinous differentiation "
"with poor differentiation/medullary architecture in highly salient hotspots."
)
interpretation = (
"The findings are strongly suggestive of a hypermutated MSI-High clinical phenotype. Colorectal carcinomas with "
"this phenotype often exhibit mismatch repair deficiency (dMMR) and typically demonstrate favorable responsiveness "
"to immune checkpoint inhibitors (PD-1/PD-L1 blockade)."
)
else:
morphology = (
"Saliency map highlights preserved, well-formed glandular structures with typical elongated, pseudostratified nuclei "
"along the crypt baselines. No significant intratumoral lymphocytic infiltration is localized within key diagnostic hotspots."
)
interpretation = (
"The morphological and cross-modal embeddings indicate a Microsatellite Stable (MSS) genotype. Standard adjuvant chemotherapy "
"regimens or fluorouracil-based therapies remain the therapeutic baseline, as responsiveness to single-agent immunotherapy is generally restricted in MSS tumors."
)
# Build references to retrieved cases
matched_cases_str = ", ".join([f"Patient {c['patient_id']} ({c['status']}, similarity: {c['similarity']:.1%})" for c in retrieved_cases])
confidence = "High (concordant model prediction and clinical retrieval)" if len(retrieved_cases) > 0 and retrieved_cases[0]["status"] == ("MSI-High" if is_msi_high else "MSS") else "Moderate"
report = f"""**PATHOLOGY REPORT: MULTIMODAL MSI PREDICTION**
**1. Morphological Observations:**
{morphology}
**2. Correlation with Historical Cohort:**
The query patch exhibits high embedding alignment with historical reference cases: {matched_cases_str}. These matching cases present similar histomorphological patterns of {retrieved_cases[0]['visual_description'] if len(retrieved_cases) > 0 else 'glandular organization'}.
**3. Clinical Interpretation:**
With a predicted **{probability:.2%}** probability of MSI-High phenotype, this tissue specimen is classified as **{status_str}**. {interpretation}
**4. Confidence and Limitations:**
Confidence Level: **{confidence}**.
*Limitations:* This assessment is restricted to a single 224x224 pixel H&E tissue patch. Formal clinical diagnosis requires full-slide histopathological correlation, MMR immunohistochemistry (IHC) for MLH1/MSH2/MSH6/PMS2, or gold-standard PCR/NGS sequencing.
*Report synthesized via GemmaSight MedGemma Pathology Assistant (Simulation).*"""
return report
|