blackmistcode's picture
Presupuesto de tokens menor con esquema (la restriccion es mas lenta por token)
30c3e85 verified
Raw
History Blame Contribute Delete
11.8 kB
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('<unused95>')
EOT_ID = processor.tokenizer.convert_tokens_to_ids('<end_of_turn>')
# 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 <unused94>thought ... <unused95> respuesta, y extract_response se queda sólo con lo
# posterior a <unused95>. 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
# <unused95>) 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 <unused95> 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 <unused95> 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 <unused95> 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()