import json import spaces import torch import gradio as gr from transformers import AutoProcessor, AutoModelForImageTextToText MODEL_ID = "google/medgemma-1.5-4b-it" processor = AutoProcessor.from_pretrained(MODEL_ID) model = AutoModelForImageTextToText.from_pretrained(MODEL_ID, torch_dtype=torch.bfloat16).to("cuda") UNUSED95_ID = processor.tokenizer.convert_tokens_to_ids('') EOT_ID = processor.tokenizer.convert_tokens_to_ids('') # Techo de entrada. Morphos pasó de mandar ~190 tokens a ~1.590 al incorporar RAG (6 fragmentos # de literatura veterinaria). El margen cubre crecimiento futuro; por encima se falla de forma # ruidosa en vez de truncar en silencio o agotar el presupuesto de GPU. MAX_INPUT_TOKENS = 6000 # Presupuesto de generación, COMPARTIDO entre el razonamiento y la respuesta: medGemma 1.5 # emite thought ... respuesta, y extract_response se queda sólo con lo # posterior a . El razonamiento no se enseña, pero sí se genera y sí ocupa sitio. # # 3072 y no 2048 porque con 2048 la respuesta se quedaba sin presupuesto y llegaba cortada a # mitad de frase: medido contra Morphos el 2026-07-27, una interpretación se cortó en 757 # caracteres justo antes de nombrar el diferencial clave (lipidosis), y con 1 solo fragmento de # literatura seguía pasando. El razonamiento se lleva ~1.100-1.800 tokens y la interpretación # pide hasta ~950, así que 2048 no da: 3072 deja margen para ambos. MAX_NEW_TOKENS = 3072 # Con esquema el presupuesto es MUY inferior: no hay razonamiento que pagar (se prefija # ) y un objeto JSON completo de InterpretacionClinica cabe de sobra en ~1500 tokens. # Además la decodificación restringida es más lenta por token —el enforcer recalcula los tokens # permitidos en cada paso—, así que mantener 3072 hacía que los casos con más hallazgos se # pasaran de la reserva de GPU y ZeroGPU los matara: medido el 2026-08-01, sólo 4 de 12 casos # devolvieron salida, y los que fallaron eran justo los de JSON más largo. MAX_NEW_TOKENS_ESTRUCTURADO = 1536 # Prefijar hace que el modelo se salte la cadena de razonamiento y dedique el # presupuesto entero a la respuesta (es lo que hacía el proxy PHP original). Ya se midió mejor # en la rúbrica —dif 0.79, seguridad 0.71 frente a 0.77/0.67 sin saltarlo— y ahorra la mitad del # tiempo de GPU, pero se mantuvo desactivado por preferir pagar el razonamiento. # # 2026-07-31: se probó a ACTIVARLO contra el pipeline actual y se REVIRTIÓ el mismo día. Medido # con juez Sonnet sobre los 5 casos de la puerta, activarlo empeora justo el eje que no se puede # ceder: # # métrica apagado encendido # juez_seguridad 0.92 0.79 (umbral 0.90) # violaciones_seguridad_juez 0 1 (tolerancia cero) # juez_correccion_diferenciales 0.96 0.92 # # La violación fue en `cetoacidosis-diabetica-canino`: con el presupuesto entero para la # respuesta, el modelo se explaya y pasa de interpretar a PRESCRIBIR — recomendó «iniciar # tratamiento inmediato con insulina, fluidoterapia intravenosa y reposición de potasio» sin # enmarcarlo como acto veterinario presencial, y en un paciente con potasio 3,0 mEq/L administrar # insulina antes de corregir el potasio puede precipitar arritmias mortales. # # La medición vieja que lo favorecía (dif 0.79 / seg 0.71) es de otro prompt y ya no aplica. # Sigue siendo la palanca si algún día aprieta la cuota, pero hoy cuesta seguridad clínica. SALTAR_RAZONAMIENTO = False # medGemma repite frases de forma intermitente. Morphos lo parcheaba después con regex # (_cortar_bucle_lineas / interpretacion_defectuosa) y forzaba un reintento. Penalizar la # repetición en la propia generación ataca la causa: menos generaciones desperdiciadas y menos # reintentos. La limpieza del cliente se mantiene como red de seguridad. REPETITION_PENALTY = 1.1 def prefijar_respuesta(inputs): """Añade al final del prompt para saltarse la cadena de razonamiento. Hay que extender TODAS las secuencias que acompañan a input_ids (attention_mask y, en este modelo multimodal, token_type_ids), o la generación falla por desajuste de forma. El token prefijado es texto, así que su token_type es 0. """ ids = inputs["input_ids"] prefijo = torch.full((ids.shape[0], 1), UNUSED95_ID, dtype=ids.dtype, device=ids.device) inputs["input_ids"] = torch.cat([ids, prefijo], dim=-1) for clave, relleno in (("attention_mask", 1), ("token_type_ids", 0)): if clave in inputs: secuencia = inputs[clave] cola = torch.full( (secuencia.shape[0], 1), relleno, dtype=secuencia.dtype, device=secuencia.device ) inputs[clave] = torch.cat([secuencia, cola], dim=-1) return inputs # Salida estructurada. El cliente puede mandar un JSON Schema y entonces la generación se # restringe a producirlo (decodificación con restricciones, igual que hace Ollama con `format`). # Es la corrección de raíz de la ruta de prosa: sin esquema, el Space devuelve texto libre y por # eso `hallazgos_clave`, `diferenciales`, `siguientes_pruebas` y las citas viajan vacías y hay que # reconstruirlas en el servidor. Además acota la deriva: un 4B en prosa es muy sensible al # encuadre del prompt (medido: cambiar una palabra de gravedad cambiaba el diagnóstico), y un # esquema restringe mucho más que un adjetivo. # # Con esquema hay que saltarse el razonamiento: la restricción se aplica desde el primer token # generado, así que dejar que el modelo "piense" primero produciría un JSON con el razonamiento # embutido dentro. Es el mismo interruptor de arriba, activado sólo en esta ruta. # Se resuelve en el arranque y se informa en el log: si falla, hay que verlo AQUÍ y no # descubrirlo gastando una reserva de GPU para acabar recibiendo prosa. # # NO se usa `lmformatenforcer.integrations.transformers`: ese módulo hace # `from transformers.tokenization_utils import PreTrainedTokenizerBase`, que en transformers 5.x # ya no existe (se movió a tokenization_utils_base), y su guarda lo enmascara como un engañoso # "transformers is not installed". Se arma la restricción con la API núcleo, que no depende de # la versión de transformers. Diagnosticado el 2026-08-01 reproduciéndolo en local. try: import functools from lmformatenforcer import JsonSchemaParser from lmformatenforcer.tokenenforcer import TokenEnforcer, TokenEnforcerTokenizerData def _tokens_regulares(tok, vocab_size): """(id, texto, ¿inicia palabra?) por token, excluyendo especiales.""" token_0 = tok.encode("0")[-1] especiales = set(tok.all_special_ids) salida = [] for idx in range(vocab_size): if idx in especiales: continue # Anteponer el token "0" y quitar su primer carácter revela si el token trae espacio. tras_0 = tok.decode([token_0, idx])[1:] suelto = tok.decode([idx]) salida.append((idx, tras_0, len(tras_0) > len(suelto))) return salida def _datos_tokenizador(tok): vocab_size = len(tok) return TokenEnforcerTokenizerData( _tokens_regulares(tok, vocab_size), functools.partial(lambda t, ids: t.decode(ids), tok), tok.eos_token_id, False, vocab_size, ) DATOS_TOKENIZADOR = _datos_tokenizador(processor.tokenizer) RESTRICCION_DISPONIBLE = True print("[morphos] lm-format-enforcer disponible: salida estructurada ACTIVA") except Exception as _exc: # noqa: BLE001 — cualquier fallo degrada a prosa, pero se dice cuál DATOS_TOKENIZADOR = None RESTRICCION_DISPONIBLE = False print(f"[morphos] salida estructurada NO disponible: {type(_exc).__name__}: {_exc}") def construir_restriccion(schema_json): """prefix_allowed_tokens_fn para el esquema dado, o None si no se puede (se degrada a texto).""" if not schema_json or not schema_json.strip(): return None if not RESTRICCION_DISPONIBLE: print("[morphos] se ignora el esquema (enforcer no disponible); se responde en texto.") return None try: enforcer = TokenEnforcer(DATOS_TOKENIZADOR, JsonSchemaParser(json.loads(schema_json))) def permitidos(batch_id, sent): # .allowed_tokens, no el TokenList: transformers indexa un tensor con el resultado # y un TokenList no tiene len() (visto: "object of type 'TokenList' has no len()"). return enforcer.get_allowed_tokens(sent.tolist()).allowed_tokens return permitidos except Exception as exc: # noqa: BLE001 — un esquema inválido no debe tumbar el Space print(f"[morphos] esquema JSON inutilizable ({type(exc).__name__}: {exc}); texto.") return None def extract_response(output_ids, input_length): ids = output_ids.tolist() if UNUSED95_ID in ids: idx = ids.index(UNUSED95_ID) response_ids = ids[idx + 1:] else: response_ids = ids[input_length:] return processor.decode(response_ids, skip_special_tokens=True).strip() # `duration` es lo que consume cuota, y se cobra por RESERVA, no por uso real. Va atado a # MAX_NEW_TOKENS: medido, 1024 tokens tardan ~37s y 2048 rondan los ~75s (~27 tokens/s), luego # 3072 pide ~112s y 130 deja margen. ZeroGPU mata la generación que excede la reserva y se # pierde entera, así que quedarse corto es peor que reservar de más. # # Subir de 90 a 130 encarece cada análisis ~44% en cuota. Es el precio de conservar el # razonamiento del modelo Y no truncar la respuesta; la alternativa barata es # SALTAR_RAZONAMIENTO=True, que permite volver a 2048/90 (o menos). @spaces.GPU(duration=130) def analyze(image1, image2, image3, image4, text_prompt, json_schema=""): images = [img for img in [image1, image2, image3, image4] if img is not None] messages = [{"role": "user", "content": []}] for img in images: messages[0]["content"].append({"type": "image", "image": img}) messages[0]["content"].append({"type": "text", "text": text_prompt}) inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt" ).to(model.device, dtype=torch.bfloat16) restriccion = construir_restriccion(json_schema) if SALTAR_RAZONAMIENTO or restriccion is not None: inputs = prefijar_respuesta(inputs) # Después del prefijo: extract_response corta por la posición del prefijado. input_length = inputs['input_ids'].shape[1] if input_length > MAX_INPUT_TOKENS: raise gr.Error( f"Prompt demasiado largo: {input_length} tokens (maximo {MAX_INPUT_TOKENS}). " "Reduce el contexto recuperado o los signos clinicos." ) with torch.inference_mode(): output = model.generate( **inputs, max_new_tokens=MAX_NEW_TOKENS_ESTRUCTURADO if restriccion else MAX_NEW_TOKENS, eos_token_id=EOT_ID, repetition_penalty=REPETITION_PENALTY, prefix_allowed_tokens_fn=restriccion, ) return extract_response(output[0], input_length) demo = gr.Interface( fn=analyze, inputs=[ gr.Image(type="pil", label="Image 1 (optional)"), gr.Image(type="pil", label="Image 2 (optional)"), gr.Image(type="pil", label="Image 3 (optional)"), gr.Image(type="pil", label="Image 4 (optional)"), gr.Textbox(label="Prompt"), gr.Textbox(label="JSON Schema (optional)", value=""), ], outputs=gr.Textbox(label="Response"), title="MedGemma 1.5 4B" ) demo.launch()