import base64 import json import re import time import boto3 # Some Gemma checkpoints wrap an internal reasoning/scratchpad span in these # special-vocab tokens before the real answer; decoded as literal text here # instead of being suppressed, so it must be stripped before display. _THINKING_RE = re.compile(r".*?", re.DOTALL) REGION = "us-east-1" BUCKET = "medgemma-async-2026" ENDPOINT = "medgemma-vllm-endpoint" _sagemaker = boto3.client("sagemaker-runtime", region_name=REGION) _s3 = boto3.client("s3", region_name=REGION) def encode_image_b64(image_bytes: bytes) -> str: return base64.b64encode(image_bytes).decode() def build_payload(messages: list[dict], image_b64: str | None, max_new_tokens: int = 300, temperature: float = 0.2) -> dict: """DJL-LMI's chat-completions path (djl_python/chat_completions) is triggered by a top-level "messages" key and expects OpenAI-style params (max_tokens, temperature) at the top level too -- NOT nested under "parameters", and NOT under an "inputs" key (that key is reserved for the plain text-generation handler, which rejects lists). """ if image_b64: last = messages[-1] last["content"] = [ {"type": "text", "text": last["content"]}, {"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image_b64}"}}, ] return { "messages": messages, "max_tokens": max_new_tokens, "temperature": temperature, } class InferenceTimeout(Exception): pass def invoke_and_wait(payload: dict, timeout_s: int = 1200, poll_interval_s: int = 5) -> dict: """Submit an async inference job and poll S3 for the result. Cold start (scale-to-zero -> instance launch -> model load) can take 10-15+ minutes on ml.g5.xlarge; timeout_s defaults high to cover that. """ input_key = f"input/chat_{int(time.time() * 1000)}.json" _s3.put_object(Bucket=BUCKET, Key=input_key, Body=json.dumps(payload).encode()) response = _sagemaker.invoke_endpoint_async( EndpointName=ENDPOINT, InputLocation=f"s3://{BUCKET}/{input_key}", ContentType="application/json", ) output_key = response["OutputLocation"].replace(f"s3://{BUCKET}/", "") failure_key = output_key.replace("output/", "failures/") deadline = time.time() + timeout_s while time.time() < deadline: try: obj = _s3.get_object(Bucket=BUCKET, Key=output_key) return json.loads(obj["Body"].read()) except _s3.exceptions.NoSuchKey: pass try: obj = _s3.get_object(Bucket=BUCKET, Key=failure_key) raise RuntimeError(f"Inference failed: {obj['Body'].read().decode()}") except _s3.exceptions.NoSuchKey: pass time.sleep(poll_interval_s) raise InferenceTimeout(f"No result after {timeout_s}s — endpoint may still be cold-starting") def extract_answer(result: dict) -> str: truncated = False if "choices" in result: choice = result["choices"][0] text = choice["message"]["content"] truncated = choice.get("finish_reason") == "length" elif "generated_text" in result: text = result["generated_text"] else: return json.dumps(result) if "" in text and "" not in text: # Generation was cut off mid-reasoning before the real answer ever # started -- nothing usable to show, so say so rather than dump the # raw scratchpad. return ("The model ran out of its response budget while reasoning " "and never reached an answer. Try again, or ask a more " "specific question.") answer = _THINKING_RE.sub("", text).strip() if truncated: answer += ("\n\n*(Response was cut off — it hit the token limit " "before finishing. Try asking a more specific question, " "or ask it to continue.)*") return answer