wuff-mann commited on
Commit
304d804
·
verified ·
1 Parent(s): 8e79cd7

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +142 -114
app.py CHANGED
@@ -1,144 +1,172 @@
1
- import os
2
- import time
3
  import gradio as gr
 
 
 
 
 
4
  from huggingface_hub import InferenceClient
5
- from mann_engram_en.router import MANNEngramRouter
6
 
7
- # 环境配置
8
- os.environ["HF_HOME"] = "/home/user/.cache/huggingface"
9
- router_engine = None
 
 
 
 
 
 
10
 
11
- def init_engine():
12
- global router_engine
13
- if router_engine is None:
14
- print("Initializing SiGLIP Tensor Routing Core...")
15
- router_engine = MANNEngramRouter(
16
- ckpt_path="./weights/skew_model_v4full_en.pt",
17
- enable_local_intent=False
18
- )
19
- return router_engine
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
20
 
21
- def extract_intent_via_api(messy_text, hf_token):
22
- if not hf_token: return "ERROR: MISSING TOKEN"
23
  client = InferenceClient("Qwen/Qwen2.5-72B-Instruct", token=hf_token)
24
- sys_prompt = (
25
- "You are an expert clinical triage AI. Extract ONLY the critical medical symptoms and targeted body parts. "
26
- "IGNORE all noise. Output a single concise English sentence."
27
- )
 
 
 
 
 
 
 
 
28
  try:
29
- res = client.chat_completion(
30
- messages=[{"role": "system", "content": sys_prompt},{"role": "user", "content": messy_text}],
31
- max_tokens=60
32
- )
33
- return res.choices[0].message.content.strip()
 
 
 
34
  except Exception as e:
35
- return f"API ERROR: {str(e)}"
36
 
37
- def process_input(query, files, hf_token, top_p_val):
38
- if not hf_token:
39
- return "❌ Error", "Missing Token", "Status: Failed.", []
40
- if not query:
41
- return "⚠️ Warning", "Empty Query", "Status: Waiting.", []
42
-
43
- engine = init_engine()
44
-
45
- file_paths = []
46
- image_paths = []
47
- if files:
48
- for f in files:
49
- ext = f.name.lower().split('.')[-1]
50
- if ext in ['jpg', 'jpeg', 'png', 'bmp', 'webp']:
51
- image_paths.append(f.name)
52
- else:
53
- file_paths.append(f.name)
54
 
 
 
 
 
55
  start_time = time.time()
 
 
 
56
 
57
- # 1. 云端解析意图
58
- clean_intent = extract_intent_via_api(query, hf_token)
59
- if "API ERROR" in clean_intent:
60
- return "Cloud API Error", clean_intent, "Check Token.", []
61
-
62
- # 2. 边缘路由 (使用动态调节的 top_p_val)
63
- results = engine.compress(
64
- query=clean_intent,
65
- context_pool=[],
66
- image_pool=image_paths,
67
- top_p=top_p_val # <--- 核心改动:应用滑块数值
68
- )
69
 
70
- latency = time.time() - start_time
71
- stats_dict = results.get("stats", {})
 
 
 
72
 
73
- stats_report = (
74
- f"✅ Routing Complete\n"
75
- f"---------------------------\n"
76
- f"🎯 Active Top_p: {top_p_val}\n"
77
- f"🧠 Cloud Brain: Qwen-72B\n"
78
- f"🖼️ Images Retained: {stats_dict.get('retained_images', 0)} / {stats_dict.get('original_images', 0)}\n"
79
- f"⚡ Total Latency: {latency:.2f}s"
80
- )
81
 
82
- return clean_intent, "", stats_report, results.get("purified_images_pil", [])
83
 
84
  # ==========================================
85
- # UI 布局:加入 Slider 控制器
86
  # ==========================================
87
- with gr.Blocks(title="MANN-Engram Router", css=".gradio-container {max-width: 100% !important;}") as demo:
88
-
89
- gr.Markdown("# 🧠 MANN-Engram: Edge-Cloud Multimodal Router")
90
 
91
  with gr.Row():
92
- # --- 左侧:配置与参数调节 ---
93
- with gr.Column(scale=2, min_width=250):
94
- gr.Markdown("### ⚙️ Settings")
95
- token_input = gr.Textbox(
96
- label="Hugging Face API Token",
97
- placeholder="hf_xxxxxxxxxxxxxxxxx",
98
- type="password"
99
- )
100
 
101
- # 💡 核心新增:Top_p 调节滑块
102
- top_p_slider = gr.Slider(
103
- minimum=0.1,
104
- maximum=1.0,
105
- value=0.85,
106
- step=0.05,
107
- label="Routing Threshold (Top_p)",
108
- info="Lower = Stricter Filtering (Precision); Higher = More Context (Recall)"
109
  )
110
 
111
- gr.Markdown("---")
112
- gr.Markdown(
113
- "**Quick Tip:**\n"
114
- "To get *ONLY* the brain tumor in Case 1, try setting Top_p to **0.5 or 0.6**."
115
  )
116
-
117
- # --- 中间:输入 ---
118
- with gr.Column(scale=4):
119
- gr.Markdown("### 📥 Input Console")
120
- query_input = gr.Textbox(label="Clinical Complaint", lines=8)
121
- file_input = gr.File(label="Patient Data Dump", file_count="multiple")
122
- submit_btn = gr.Button("🚀 Execute Routing", variant="primary", size="lg")
123
 
124
- gr.Examples(
125
- examples=[["Doctor, I have severe stomach pain, but my real concern is the seizure I had today and the numbness in my left arm. Check my head MRI.", None]],
126
- inputs=[query_input, file_input],
127
- label="Preset Scenario"
128
- )
129
 
130
- # --- 右侧:结果 ---
131
- with gr.Column(scale=4):
132
- gr.Markdown("### 📤 Output Dashboard")
133
- out_intent = gr.Textbox(label="🎯 Extracted Clinical Intent", interactive=False)
134
- out_stats = gr.Textbox(label="📊 Metrics", interactive=False, lines=6)
135
- out_gallery = gr.Gallery(label="🔭 Routed Evidence", columns=2, height="auto")
 
 
 
 
136
 
137
- # 点击事件绑定,加入 top_p_slider 输入
138
- submit_btn.click(
139
- fn=process_input,
140
- inputs=[query_input, file_input, token_input, top_p_slider],
141
- outputs=[out_intent, gr.State(), out_stats, out_gallery]
142
  )
143
 
144
  if __name__ == "__main__":
 
 
 
1
  import gradio as gr
2
+ import torch
3
+ import json
4
+ import time
5
+ from PIL import Image
6
+ from transformers import AutoProcessor, AutoModel
7
  from huggingface_hub import InferenceClient
 
8
 
9
+ # ==========================================
10
+ # Phase 2: Edge-Side Tensor Router (SiGLIP)
11
+ # ==========================================
12
+ class SiGLIPRouter:
13
+ def __init__(self):
14
+ print("Loading Edge Routing Engine (SiGLIP-So400M)...")
15
+ # Load locally for edge-simulated routing
16
+ self.processor = AutoProcessor.from_pretrained("google/siglip-so400m-patch14-384")
17
+ self.model = AutoModel.from_pretrained("google/siglip-so400m-patch14-384")
18
 
19
+ def route_evidence(self, visual_query, image_paths, margin_ratio=0.15, absolute_floor=0.1):
20
+ """
21
+ Relative Margin Thresholding Engine
22
+ Routes images based on dynamic confidence window relative to the best match.
23
+ """
24
+ if not image_paths:
25
+ return [], {}
26
+
27
+ # Convert file paths to PIL Images
28
+ images = [Image.open(img).convert("RGB") for img in image_paths]
29
+
30
+ # Extract multimodal tensors
31
+ inputs = self.processor(text=[visual_query], images=images, padding="max_length", return_tensors="pt")
32
+
33
+ with torch.no_grad():
34
+ outputs = self.model(**inputs)
35
+
36
+ # Get Sigmoid matching probabilities
37
+ logits_per_image = outputs.logits_per_image
38
+ probs = torch.sigmoid(logits_per_image).squeeze()
39
+
40
+ # Handle single image fallback
41
+ if probs.dim() == 0:
42
+ probs = probs.unsqueeze(0)
43
+
44
+ # Calculate dynamic threshold based on the Anchor (Highest Prob)
45
+ max_prob = torch.max(probs).item()
46
+ dynamic_threshold = max(max_prob * (1.0 - margin_ratio), absolute_floor)
47
+
48
+ # Precision Pruning
49
+ routed_images = []
50
+ for idx, prob in enumerate(probs):
51
+ if prob.item() >= dynamic_threshold:
52
+ routed_images.append(image_paths[idx])
53
+
54
+ metrics = {
55
+ "Anchor_Probability (Max)": round(max_prob, 4),
56
+ "Dynamic_Threshold": round(dynamic_threshold, 4),
57
+ "Total_Candidates": len(image_paths),
58
+ "Images_Retained": len(routed_images)
59
+ }
60
+
61
+ return routed_images, metrics
62
+
63
+ # ==========================================
64
+ # Phase 1: Cloud-Side Distillation (Qwen)
65
+ # ==========================================
66
+ def extract_intent_and_query(clinical_text, hf_token):
67
+ """
68
+ Dual-track extraction using Qwen-72B via Hugging Face Inference API.
69
+ Translates messy text into Clinical Intent and Visual Target.
70
+ """
71
+ if not hf_token:
72
+ return {"error": "Missing Hugging Face API Token."}, "Error"
73
 
 
 
74
  client = InferenceClient("Qwen/Qwen2.5-72B-Instruct", token=hf_token)
75
+
76
+ system_prompt = """You are an expert Clinical Triage AI and a Multimodal Routing Specialist.
77
+ Your task is to analyze messy, noisy patient narratives and output a strictly formatted JSON object with two fields.
78
+ 1. "Clinical_Intent": A purified medical summary of the core issue.
79
+ 2. "Visual_Query": An extreme extraction of visual, anatomical, and radiological keywords relevant ONLY to the core issue. Think like an image-recognition model. Use nouns and modalities (e.g., 'Brain MRA, skull, Willis circle'). DO NOT use abstract symptoms like 'headache'.
80
+ Output ONLY valid JSON."""
81
+
82
+ messages = [
83
+ {"role": "system", "content": system_prompt},
84
+ {"role": "user", "content": f"Patient Narrative:\n{clinical_text}"}
85
+ ]
86
+
87
  try:
88
+ response = client.chat_completion(messages=messages, max_tokens=300)
89
+ content = response.choices[0].message.content
90
+
91
+ # Clean markdown formatting if present
92
+ clean_content = content.replace("```json", "").replace("```", "").strip()
93
+ result = json.loads(clean_content)
94
+
95
+ return result, result.get("Visual_Query", "medical scan")
96
  except Exception as e:
97
+ return {"error": f"Cloud Distillation Failed: {str(e)}"}, "medical scan"
98
 
99
+ # Initialize local router
100
+ router = SiGLIPRouter()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
101
 
102
+ # ==========================================
103
+ # Main Execution Pipeline
104
+ # ==========================================
105
+ def execute_pipeline(hf_token, narrative, images, margin_ratio):
106
  start_time = time.time()
107
+
108
+ if not images:
109
+ return {"Error": "No images uploaded."}, [], {"Status": "Failed"}
110
 
111
+ # 1. Cloud Intelligence (Linguistic Distillation)
112
+ cloud_result, visual_query = extract_intent_and_query(narrative, hf_token)
 
 
 
 
 
 
 
 
 
 
113
 
114
+ if "error" in cloud_result:
115
+ return cloud_result, [], {"Status": "Cloud API Error"}
116
+
117
+ # 2. Edge Intelligence (Tensor Routing)
118
+ routed_imgs, metrics = router.route_evidence(visual_query, images, margin_ratio=margin_ratio)
119
 
120
+ end_time = time.time()
121
+ metrics["Total_Latency (s)"] = round(end_time - start_time, 2)
122
+ metrics["Active_Margin_Ratio"] = margin_ratio
 
 
 
 
 
123
 
124
+ return cloud_result, routed_imgs, metrics
125
 
126
  # ==========================================
127
+ # Gradio UI Design
128
  # ==========================================
129
+ with gr.Blocks(title="MANN-Engram Showcase", theme=gr.themes.Base()) as demo:
130
+ gr.Markdown("# 🧠 MANN-Engram: Edge-Cloud Multimodal Semantic Router")
131
+ gr.Markdown("> **A Privacy-First, Zero-Hallucination Shield for Clinical Vision-Language Models.**")
132
 
133
  with gr.Row():
134
+ # Left Column: Inputs & Settings
135
+ with gr.Column(scale=1):
136
+ gr.Markdown("### ⚙️ Engine Settings")
137
+ hf_token = gr.Textbox(label="Hugging Face API Token (For Cloud Brain)", type="password", placeholder="hf_xxxxxxxx...")
 
 
 
 
138
 
139
+ margin_slider = gr.Slider(
140
+ minimum=0.05, maximum=0.40, step=0.05, value=0.15,
141
+ label="Routing Tolerance (Margin Ratio)",
142
+ info="Lower (0.05) = Sniper Mode (Extreme Precision). Higher (0.30) = Cluster Mode (Recalls multiple related views)."
 
 
 
 
143
  )
144
 
145
+ gr.Markdown("### 📥 Patient Data Dump")
146
+ narrative_input = gr.Textbox(
147
+ label="Messy Clinical Narrative", lines=8,
148
+ placeholder="Paste the chaotic patient complaint and history here..."
149
  )
150
+ image_input = gr.File(label="Upload Unorganized Scans (Images)", file_count="multiple", type="filepath")
 
 
 
 
 
 
151
 
152
+ run_btn = gr.Button("🚀 Execute Routing Pipeline", variant="primary")
 
 
 
 
153
 
154
+ # Right Column: Outputs
155
+ with gr.Column(scale=1):
156
+ gr.Markdown("### ☁️ Cloud Output: Dual-Track Distillation")
157
+ cloud_output = gr.JSON(label="Purified Intent & Visual Query")
158
+
159
+ gr.Markdown("### 🛡️ Edge Output: Routed Core Evidence")
160
+ routed_gallery = gr.Gallery(label="Surgically Selected Scans", columns=2, object_fit="contain", height=400)
161
+
162
+ gr.Markdown("### 📊 Telemetry & Metrics")
163
+ metrics_output = gr.JSON(label="Routing Diagnostics")
164
 
165
+ # Wire up the button
166
+ run_btn.click(
167
+ fn=execute_pipeline,
168
+ inputs=[hf_token, narrative_input, image_input, margin_slider],
169
+ outputs=[cloud_output, routed_gallery, metrics_output]
170
  )
171
 
172
  if __name__ == "__main__":