Image Hopper router v2-fast: faster serving and default-route image cap (same weights and calibration)
0048827 verified Download open_decisions/image_jev/vision.py from HopitAI/image-hopper: direct link, hf CLI and curl.
- Browser
- Download file 28.8 kB
-
https://huggingface.co/HopitAI/image-hopper/resolve/main/open_decisions/image_jev/vision.py
- Command line
-
hf download hf://HopitAI/image-hopper/open_decisions/image_jev/vision.py
-
curl -L -o vision.py https://huggingface.co/HopitAI/image-hopper/resolve/main/open_decisions/image_jev/vision.py
28.8 kB
| """Shared image preparation and option readout for Image JevBench. | |
| This module is deliberately the only path from source images to processor inputs. Training and | |
| serving should both call :func:`prepare_for_route`; keeping resize, prompt and readout decisions | |
| here makes a route or budget change visible and versioned instead of an accidental processor | |
| configuration change. | |
| """ | |
| from __future__ import annotations | |
| import base64 | |
| import binascii | |
| import io | |
| import json | |
| import hashlib | |
| import inspect | |
| import math | |
| import os | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from types import MappingProxyType | |
| from typing import Any | |
| from PIL import Image, ImageOps, UnidentifiedImageError | |
| from open_decisions.image_jev.speed import PhaseTimer, move_to_device | |
| # One source for option rendering in training and serving (no fallback path that could drift); | |
| # tests/test_image_jev_vision.py checks it against the public hopper_decisions copy when that is importable. | |
| from open_decisions.scoring import prompt as hopper_prompt | |
| FACTOR = 32 | |
| MIN_PIXELS = 65_536 | |
| MAX_PIXELS = 16_777_216 | |
| MAX_IMAGE_BYTES = 20 * 1024 * 1024 | |
| MAX_IMAGE_PIXELS = 20_000_000 | |
| # Optional native-route image-token cap (a serving option, off by default). The cap lowers the | |
| # processor's own max pixels to ``cap * FACTOR**2`` so its smart resize keeps at most ``cap`` | |
| # merged image tokens per image. Below 256 tokens the processor's minimum-area rule and its | |
| # one-factor edge clamp can exceed the cap for extreme aspect ratios, so smaller caps are refused. | |
| MIN_IMAGE_TOKEN_CAP = 256 | |
| MAX_IMAGE_TOKEN_CAP = MAX_PIXELS // (FACTOR * FACTOR) | |
| BUDGET_VERSION = "b1" | |
| BUDGETS = MappingProxyType({"photo": 160, "scene": 224, "document": 384, "dense": 768}) | |
| ROUTING_VERSION = "image-jev/content-routing/v1" | |
| ROUTING_CONFIG = MappingProxyType({ | |
| "version": ROUTING_VERSION, | |
| "inputs": "decoded-oriented-rgb-pixels-only", | |
| "policy": "conservative-document", | |
| "content_class": "document", | |
| }) | |
| ROUTING_CONFIG_SHA256 = hashlib.sha256( | |
| json.dumps(dict(ROUTING_CONFIG), sort_keys=True, separators=(",", ":")).encode("utf-8") | |
| ).hexdigest() | |
| IMAGE_ROUTE_VERSION = "image-jev/image-route/v1" | |
| IMAGE_ROUTES = ("native", "capped") | |
| IMAGE_ROUTE_CONFIG_SHA256 = MappingProxyType({ | |
| route: hashlib.sha256(json.dumps( | |
| {"version": IMAGE_ROUTE_VERSION, "image_route": route}, | |
| sort_keys=True, separators=(",", ":"), | |
| ).encode("utf-8")).hexdigest() | |
| for route in IMAGE_ROUTES | |
| }) | |
| IMAGE_SYSTEM = ("You make decisions about supplied images. Inspect every image and reply with the letter of the " | |
| "correct option and nothing else. Do not produce reasoning.") | |
| VISION_TOKEN = "<|vision_start|><|image_pad|><|vision_end|>" | |
| class PreparedInput: | |
| """The complete model input produced identically for training and serving.""" | |
| input_ids: Any | |
| pixel_values: Any | |
| image_grid_thw: Any | |
| resize_records: list[dict[str, Any]] | |
| letter_token_ids: list[int] | |
| prompt: str | |
| routing_config_sha256: str = ROUTING_CONFIG_SHA256 | |
| # transformers >= 5.x Qwen-VL needs this for multimodal RoPE; returned by the processor beside input_ids. | |
| mm_token_type_ids: Any = None | |
| image_route: str = "capped" | |
| def ids(self): | |
| """Short alias used by call sites that name language inputs ``ids``.""" | |
| return self.input_ids | |
| def pixels(self): | |
| """Short alias for the processor's image tensor.""" | |
| return self.pixel_values | |
| def records(self): | |
| """Short alias for the resize audit records.""" | |
| return self.resize_records | |
| def image_route_config_sha256(self): | |
| """Registered identity of the native/capped preparation choice.""" | |
| return IMAGE_ROUTE_CONFIG_SHA256[self.image_route] | |
| def smart_resize(height: int, width: int, factor: int = FACTOR, min_pixels: int = MIN_PIXELS, | |
| max_pixels: int = MAX_PIXELS) -> tuple[int, int]: | |
| """Qwen2-VL/Qwen3-VL's ``smart_resize``, with Qwen3.5-4B image defaults. | |
| The arithmetic and its boundary behavior intentionally track Transformers exactly, including | |
| its rejection only when the absolute aspect ratio is *greater than* 200. | |
| """ | |
| if max(height, width) / min(height, width) > 200: | |
| raise ValueError( | |
| f"absolute aspect ratio must be smaller than 200, got {max(height, width) / min(height, width)}" | |
| ) | |
| h_bar = round(height / factor) * factor | |
| w_bar = round(width / factor) * factor | |
| if h_bar * w_bar > max_pixels: | |
| beta = math.sqrt((height * width) / max_pixels) | |
| h_bar = max(factor, math.floor(height / beta / factor) * factor) | |
| w_bar = max(factor, math.floor(width / beta / factor) * factor) | |
| elif h_bar * w_bar < min_pixels: | |
| beta = math.sqrt(min_pixels / (height * width)) | |
| h_bar = math.ceil(height * beta / factor) * factor | |
| w_bar = math.ceil(width * beta / factor) * factor | |
| return h_bar, w_bar | |
| def image_tokens(height: int, width: int, factor: int = FACTOR) -> int: | |
| """Merged image-token count after the model processor's normal smart resize.""" | |
| resized_h, resized_w = smart_resize(height, width, factor=factor) | |
| return (resized_h // factor) * (resized_w // factor) | |
| def validate_image_token_cap(value: Any) -> int: | |
| """An explicit native-route image-token cap: an integer in [256, 16384].""" | |
| if isinstance(value, bool) or not isinstance(value, int): | |
| raise ValueError(f"image_token_cap must be an integer, got {value!r}") | |
| if not MIN_IMAGE_TOKEN_CAP <= value <= MAX_IMAGE_TOKEN_CAP: | |
| raise ValueError(f"image_token_cap must be between {MIN_IMAGE_TOKEN_CAP} and " | |
| f"{MAX_IMAGE_TOKEN_CAP}, got {value}") | |
| return value | |
| def native_max_pixels(image_token_cap: int | None = None) -> int: | |
| """The processor max pixels of the native route: its default, or ``cap * FACTOR**2``.""" | |
| if image_token_cap is None: | |
| return MAX_PIXELS | |
| return validate_image_token_cap(image_token_cap) * FACTOR * FACTOR | |
| def native_image_tokens(height: int, width: int, image_token_cap: int | None = None) -> int: | |
| """Merged image tokens of one native-route image, with the optional cap applied.""" | |
| resized_h, resized_w = smart_resize(height, width, max_pixels=native_max_pixels(image_token_cap)) | |
| return (resized_h // FACTOR) * (resized_w // FACTOR) | |
| def budget_pixels(content_class: str) -> int: | |
| """Return the explicit pixel-area budget for a named content class.""" | |
| try: | |
| return BUDGETS[content_class] * FACTOR * FACTOR | |
| except (KeyError, TypeError): | |
| raise ValueError(f"unknown image content class {content_class!r}; expected one of {tuple(BUDGETS)}") from None | |
| def route_content(img: Image.Image) -> str: | |
| """Conservatively route every image to ``document`` for now. | |
| Content routing is a later measured experiment. Any future classifier must use only observable | |
| image content, never benchmark ids, source names, family names or other hidden metadata. | |
| """ | |
| if not isinstance(img, Image.Image): | |
| raise TypeError("classify_content expects a PIL.Image.Image") | |
| return "document" | |
| # Compatibility name for callers that only need the returned class. All | |
| # consumers route through ``route_content`` and bind ROUTING_CONFIG_SHA256. | |
| classify_content = route_content | |
| def _scaled_size(width: int, height: int, long_edge: int) -> tuple[int, int]: | |
| """Largest-edge parameterisation of a whole-image, aspect-preserving integer size.""" | |
| if width >= height: | |
| return long_edge, max(1, min(height, round(height * long_edge / width))) | |
| return max(1, min(width, round(width * long_edge / height))), long_edge | |
| def _budgeted_size(width: int, height: int, token_budget: int) -> tuple[int, int]: | |
| """A whole-image size whose default processor result meets ``token_budget``.""" | |
| if image_tokens(height, width) <= token_budget: | |
| return width, height | |
| # Prefer the processor's own factor-aligned target. This makes its subsequent smart resize a | |
| # no-op and ensures LANCZOS, rather than a second processor interpolation, sets the final pixels. | |
| target_h, target_w = smart_resize(height, width, max_pixels=token_budget * FACTOR * FACTOR) | |
| while (target_h // FACTOR) * (target_w // FACTOR) > token_budget: | |
| if target_w >= target_h and target_w > FACTOR: | |
| target_w -= FACTOR | |
| elif target_h > FACTOR: | |
| target_h -= FACTOR | |
| else: | |
| break | |
| try: | |
| aligned_tokens = image_tokens(target_h, target_w) | |
| except ValueError: | |
| aligned_tokens = token_budget + 1 | |
| if target_w <= width and target_h <= height and aligned_tokens <= token_budget: | |
| return target_w, target_h | |
| # A source edge shorter than one factor can make the aligned target an upscale. Search actual | |
| # image sizes in that edge case; the saved pixels still enforce the budget even if a later | |
| # processor is constructed with its native 16M maximum. | |
| low, high = 1, max(width, height) | |
| best = None | |
| while low <= high: | |
| middle = (low + high) // 2 | |
| candidate = _scaled_size(width, height, middle) | |
| try: | |
| fits = image_tokens(candidate[1], candidate[0]) <= token_budget | |
| except ValueError: | |
| fits = False | |
| if fits: | |
| best = candidate | |
| low = middle + 1 | |
| else: | |
| high = middle - 1 | |
| if best is None: # Unreachable for accepted (<=200:1) inputs and the smallest 160-token budget. | |
| raise ValueError(f"image aspect ratio cannot fit the {token_budget}-token budget") | |
| return best | |
| def resize_for_budget(img: Image.Image, content_class: str) -> tuple[Image.Image, dict[str, Any]]: | |
| """Convert to RGB and, if necessary, LANCZOS-downscale the whole image to its named budget.""" | |
| if not isinstance(img, Image.Image): | |
| raise TypeError("resize_for_budget expects a PIL.Image.Image") | |
| token_budget = budget_pixels(content_class) // (FACTOR * FACTOR) | |
| orig_size = img.size | |
| new_size = _budgeted_size(*orig_size, token_budget) | |
| converted = img.convert("RGB") | |
| resized = converted if new_size == orig_size else converted.resize(new_size, Image.Resampling.LANCZOS) | |
| tokens = image_tokens(resized.height, resized.width) | |
| if tokens > token_budget: # Keep a hard invariant next to the operation that enforces it. | |
| raise RuntimeError(f"resize produced {tokens} image tokens for a {token_budget}-token budget") | |
| record = {"content_class": content_class, "budget_version": BUDGET_VERSION, | |
| "orig_size": orig_size, "new_size": resized.size, "image_tokens": tokens} | |
| return resized, record | |
| def _normalise_options(options) -> list[dict[str, str]]: | |
| normalised = [] | |
| for index, option in enumerate(options): | |
| if isinstance(option, str): | |
| normalised.append({"name": option, "description": option}) | |
| elif isinstance(option, dict) and "name" in option and "description" in option: | |
| normalised.append({"name": str(option["name"]), "description": str(option["description"])}) | |
| elif isinstance(option, (tuple, list)) and len(option) == 2: | |
| normalised.append({"name": str(option[0]), "description": str(option[1])}) | |
| else: | |
| raise ValueError(f"option {index} must be text, a (name, description) pair, or a mapping with those keys") | |
| if not 2 <= len(normalised) <= len(hopper_prompt.LETTERS): | |
| raise ValueError("image decisions require 2 to 26 options") | |
| return normalised | |
| def render_prompt(question: str, options, n_images: int) -> str: | |
| """Render the image decision user turn with explicit letter-description pairs and thinking off.""" | |
| if not isinstance(question, str) or not question: | |
| raise ValueError("question must be non-empty text") | |
| if not isinstance(n_images, int) or isinstance(n_images, bool) or n_images < 1: | |
| raise ValueError("n_images must be a positive integer") | |
| normalised = _normalise_options(options) | |
| shown = hopper_prompt.option_lines({"kind": "choice", "options": normalised}) | |
| payload = {"question": question, | |
| "options": [{"letter": hopper_prompt.LETTERS[i], "description": text} | |
| for i, (_, text) in enumerate(shown)]} | |
| images = "\n".join(f"IMAGE {i + 1}\n{VISION_TOKEN}" for i in range(n_images)) | |
| return (f"{images}\n\n{json.dumps(payload, ensure_ascii=False)}\n\n" | |
| "Reply with one option letter only; do not produce reasoning.") | |
| def _decode_data_uri(uri: str) -> bytes: | |
| try: | |
| header, encoded = uri.split(",", 1) | |
| except ValueError: | |
| raise ValueError("bad image data URI: missing comma") from None | |
| if not header.startswith("data:image/") or not header.endswith(";base64"): | |
| raise ValueError("bad image data URI: expected data:image/...;base64,...") | |
| # Reject an oversized payload before allocating its decoded representation. Padding means the | |
| # estimate can exceed the true size by at most two bytes, so an exact check follows decoding. | |
| if (len(encoded) // 4) * 3 > MAX_IMAGE_BYTES + 2: | |
| raise ValueError(f"image exceeds {MAX_IMAGE_BYTES} bytes") | |
| try: | |
| raw = base64.b64decode(encoded, validate=True) | |
| except (binascii.Error, ValueError): | |
| raise ValueError("bad image data URI: invalid base64 payload") from None | |
| if len(raw) > MAX_IMAGE_BYTES: | |
| raise ValueError(f"image exceeds {MAX_IMAGE_BYTES} bytes") | |
| return raw | |
| def _image_bytes(source) -> bytes: | |
| if isinstance(source, (bytes, bytearray, memoryview)): | |
| raw = bytes(source) | |
| if len(raw) > MAX_IMAGE_BYTES: | |
| raise ValueError(f"image exceeds {MAX_IMAGE_BYTES} bytes") | |
| return raw | |
| if isinstance(source, str) and source.startswith("data:"): | |
| return _decode_data_uri(source) | |
| if isinstance(source, (str, os.PathLike)): | |
| path = Path(source) | |
| try: | |
| size = path.stat().st_size | |
| except OSError as error: | |
| raise ValueError(f"cannot read image path {path}: {error}") from error | |
| if size > MAX_IMAGE_BYTES: | |
| raise ValueError(f"image exceeds {MAX_IMAGE_BYTES} bytes: {path}") | |
| try: | |
| raw = path.read_bytes() | |
| except OSError as error: | |
| raise ValueError(f"cannot read image path {path}: {error}") from error | |
| if len(raw) > MAX_IMAGE_BYTES: # The file may have changed between stat and read. | |
| raise ValueError(f"image exceeds {MAX_IMAGE_BYTES} bytes: {path}") | |
| return raw | |
| raise TypeError("each image must be bytes, a filesystem path, or a base64 image data URI") | |
| def _load_image(source, *, exif_transpose: bool = True) -> Image.Image: | |
| raw = _image_bytes(source) | |
| try: | |
| with Image.open(io.BytesIO(raw)) as opened: | |
| width, height = opened.size | |
| if width * height > MAX_IMAGE_PIXELS: | |
| raise ValueError(f"image exceeds {MAX_IMAGE_PIXELS} pixels: {width}x{height}") | |
| opened.load() | |
| ready = ImageOps.exif_transpose(opened) if exif_transpose else opened | |
| # A loaded image owns its pixels after the file closes. | |
| return ready | |
| except ValueError: | |
| raise | |
| except (UnidentifiedImageError, OSError, SyntaxError) as error: | |
| raise ValueError(f"cannot decode image: {error}") from error | |
| def _output_field(output, name: str): | |
| if isinstance(output, dict): | |
| value = output.get(name) | |
| else: | |
| value = getattr(output, name, None) | |
| if value is None: | |
| raise ValueError(f"processor output is missing {name!r}") | |
| return value | |
| def _processor_min_pixels(processor) -> int: | |
| """The processor's own minimum area (kept unchanged when a token cap lowers the maximum).""" | |
| image_processor = getattr(processor, "image_processor", None) | |
| size = getattr(image_processor, "size", None) | |
| for getter in (lambda: size["shortest_edge"], lambda: getattr(size, "shortest_edge"), | |
| lambda: getattr(image_processor, "min_pixels")): | |
| try: | |
| value = getter() | |
| except (KeyError, TypeError, AttributeError): | |
| continue | |
| if isinstance(value, int) and not isinstance(value, bool) and value > 0: | |
| return value | |
| return MIN_PIXELS | |
| def _grid_rows(grid) -> list[list[int]]: | |
| rows = grid.tolist() if hasattr(grid, "tolist") else [list(row) for row in grid] | |
| return [[int(value) for value in row] for row in rows] | |
| def _prepare(example: dict, processor, *, image_route: str, | |
| content_class: str | None = None, loaded_images=None, cache=None, | |
| timer=None, device=None, image_token_cap: int | None = None) -> PreparedInput: | |
| """Prepare one decision after validating an explicit registered image route. | |
| ``image_token_cap`` (native route only; default None = the processor's own maximum) lowers | |
| the processor's max pixels so every image keeps at most that many merged image tokens. | |
| """ | |
| if image_route not in IMAGE_ROUTES: | |
| raise ValueError(f"image_route must be one of {IMAGE_ROUTES}, got {image_route!r}") | |
| if image_token_cap is not None: | |
| if image_route != "native": | |
| raise ValueError("image_token_cap applies to the native image route only") | |
| image_token_cap = validate_image_token_cap(image_token_cap) | |
| try: | |
| question, options = example["question"], example["options"] | |
| except (KeyError, TypeError) as error: | |
| raise ValueError("example must contain question and options") from error | |
| sources = example.get("images") | |
| if sources is None and "image" in example: | |
| sources = [example["image"]] | |
| if not isinstance(sources, (list, tuple)) or not sources: | |
| raise ValueError("example must contain a non-empty images list") | |
| normalised = _normalise_options(options) | |
| if content_class is not None: | |
| raise ValueError("content_class overrides are forbidden; routing uses observable image content") | |
| # The historical native read path used decoded pixel order directly. Keep | |
| # that exact behavior while capped retains its established EXIF transpose. | |
| timer = timer or PhaseTimer() | |
| with timer.phase("decode_load"): | |
| loaded = loaded_images if loaded_images is not None else [ | |
| _load_image(source, exif_transpose=image_route == "capped") for source in sources | |
| ] | |
| with timer.phase("processor_preprocess"): | |
| ready_images, records = [], [] | |
| for image in loaded: | |
| routed_class = route_content(image) | |
| if image_route == "capped": | |
| ready, record = resize_for_budget(image, routed_class) | |
| else: | |
| ready = image if image.mode == "RGB" else image.convert("RGB") | |
| record = { | |
| "image_route": "native", | |
| "content_class": routed_class, | |
| "orig_size": image.size, | |
| "new_size": ready.size, | |
| "image_tokens": image_tokens(ready.height, ready.width), | |
| } | |
| if image_token_cap is not None: | |
| # The processor, not this function, resizes; new_size stays the decoded size. | |
| record["image_token_cap"] = image_token_cap | |
| record["image_tokens"] = native_image_tokens(ready.height, ready.width, | |
| image_token_cap) | |
| ready_images.append(ready) | |
| records.append(record) | |
| with timer.phase("tokenise_chat_template"): | |
| user = render_prompt(question, normalised, len(ready_images)) | |
| if cache is not None: | |
| prompt = cache.prompt(user) | |
| elif hasattr(processor, "apply_chat_template"): | |
| messages = [{"role": "system", "content": IMAGE_SYSTEM}, {"role": "user", "content": user}] | |
| prompt = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, | |
| enable_thinking=False) | |
| else: | |
| prompt = user | |
| tokenizer = getattr(processor, "tokenizer", None) | |
| if tokenizer is None: | |
| raise ValueError("processor must expose its tokenizer for option-letter readout") | |
| token_ids = (cache.letter_ids(len(normalised)) if cache is not None else | |
| hopper_prompt.letter_token_ids(tokenizer)[:len(normalised)]) | |
| # Fast torch image processors accept device; slow/PIL processors do not. | |
| image_processor = getattr(processor, "image_processor", None) | |
| fast = image_processor is not None and ( | |
| getattr(image_processor, "backend", None) == "torch" or | |
| "Fast" in type(image_processor).__name__ or | |
| any(base.__name__ == "BaseImageProcessorFast" for base in type(image_processor).__mro__)) | |
| images_kwargs = {"device": device} if fast and device is not None else {} | |
| if image_token_cap is not None: | |
| # Every spelling the processor's kwargs accept, all consistent: min unchanged, max capped. | |
| min_pixels, max_pixels = _processor_min_pixels(processor), native_max_pixels(image_token_cap) | |
| images_kwargs.update(size={"shortest_edge": min_pixels, "longest_edge": max_pixels}, | |
| min_pixels=min_pixels, max_pixels=max_pixels) | |
| kwargs = {"images_kwargs": images_kwargs} if images_kwargs else {} | |
| if cache is not None: | |
| cache.tokenizer.timer = timer | |
| token_before = timer.seconds["tokenise_chat_template"] if timer.seconds is not None else 0.0 | |
| with timer.phase("processor_preprocess"): | |
| output = processor(text=[prompt], images=ready_images, padding=True, return_tensors="pt", **kwargs) | |
| if image_token_cap is not None: | |
| # Fail closed if the processor ignored the cap (a library change would do that). | |
| merge = int(getattr(getattr(processor, "image_processor", None), "merge_size", 2) or 2) | |
| served = [t * h * w // (merge * merge) | |
| for t, h, w in _grid_rows(_output_field(output, "image_grid_thw"))] | |
| if len(served) != len(records) or any(tokens > image_token_cap for tokens in served): | |
| raise RuntimeError(f"processor returned {served} image tokens for an " | |
| f"{image_token_cap}-token cap") | |
| for record, tokens in zip(records, served): | |
| record["processor_image_tokens"] = tokens | |
| if timer.seconds is not None: | |
| timer.seconds["processor_preprocess"] -= timer.seconds["tokenise_chat_template"] - token_before | |
| return PreparedInput(input_ids=_output_field(output, "input_ids"), | |
| pixel_values=_output_field(output, "pixel_values"), | |
| image_grid_thw=_output_field(output, "image_grid_thw"), | |
| resize_records=records, letter_token_ids=token_ids, prompt=prompt, | |
| mm_token_type_ids=(output.get("mm_token_type_ids") if isinstance(output, dict) | |
| else getattr(output, "mm_token_type_ids", None)), | |
| image_route=image_route) | |
| def prepare(example: dict, processor, content_class: str | None = None) -> PreparedInput: | |
| """Prepare the existing capped route shared by capped training and serving.""" | |
| return _prepare(example, processor, image_route="capped", content_class=content_class) | |
| def prepare_native(example: dict, processor, **kwargs) -> PreparedInput: | |
| """Prepare processor-default full-resolution images without budget resizing. | |
| ``image_token_cap=N`` (optional, default off) lowers only the processor's max pixels. | |
| """ | |
| return _prepare(example, processor, image_route="native", **kwargs) | |
| def prepare_for_route(example: dict, processor, route: str) -> PreparedInput: | |
| """Dispatch one example through exactly one registered image route.""" | |
| if route == "native": | |
| return prepare_native(example, processor) | |
| if route == "capped": | |
| return prepare(example, processor) | |
| raise ValueError(f"image_route must be one of {IMAGE_ROUTES}, got {route!r}") | |
| def last_position_logits(model, prepared, *, timer=None, forward_call=None): | |
| """Run exactly one last-position vocabulary projection for every local reader.""" | |
| try: | |
| device = next(model.parameters()).device | |
| except (StopIteration, AttributeError): | |
| device = None | |
| timer = timer or PhaseTimer() | |
| def move(value): | |
| return move_to_device(value, device) | |
| with timer.phase("h2d_copy"): | |
| kwargs = { | |
| "input_ids": move(prepared.input_ids), | |
| "pixel_values": move(prepared.pixel_values), | |
| "image_grid_thw": move(prepared.image_grid_thw), | |
| "use_cache": False, | |
| **({"mm_token_type_ids": move(prepared.mm_token_type_ids)} | |
| if getattr(prepared, "mm_token_type_ids", None) is not None else {}), | |
| "return_dict": True, | |
| } | |
| readout_kwargs = getattr(model, "_image_last_position_kwargs", None) | |
| if readout_kwargs is None: | |
| parameters = inspect.signature(model.forward).parameters | |
| if "logits_to_keep" in parameters or any( | |
| parameter.kind == parameter.VAR_KEYWORD for parameter in parameters.values() | |
| ): | |
| readout_kwargs = {"logits_to_keep": 1} | |
| elif "num_logits_to_keep" in parameters: | |
| readout_kwargs = {"num_logits_to_keep": 1} | |
| else: | |
| raise ValueError("model does not expose a last-position logits readout") | |
| model._image_last_position_kwargs = readout_kwargs | |
| kwargs.update(readout_kwargs) | |
| with timer.phase("forward"): | |
| output = model(**kwargs) if forward_call is None else forward_call(lambda: model(**kwargs)) | |
| logits = output["logits"] if isinstance(output, dict) else output.logits | |
| if logits.ndim != 3 or logits.shape[0] != 1 or logits.shape[1] != 1: | |
| raise ValueError("model must return exactly [1, 1, vocabulary] logits") | |
| return logits[0, 0, :] | |
| def calibrated_read(model, prepared: PreparedInput, *, temperature: float) -> tuple[list[float], list[float]]: | |
| """Return raw and checkpoint-artifact-calibrated option probabilities.""" | |
| logits = last_position_logits(model, prepared) | |
| import torch | |
| raw = option_log_probs(logits, prepared.letter_token_ids, validate_finite=False).exp() | |
| calibrated = option_log_probs(logits, prepared.letter_token_ids, temperature, | |
| validate_finite=False).exp() | |
| values = torch.stack((raw, calibrated)).detach().cpu().tolist() | |
| if not all(math.isfinite(value) for row in values for value in row): | |
| raise ValueError("option probabilities are non-finite") | |
| return values[0], values[1] | |
| def option_log_probs(logits_last, letter_token_ids: list[int], temperature: float = 1.0, *, | |
| validate_finite: bool = True): | |
| """Differentiable FP32 log-softmax over the served option-letter rows. | |
| Training uses this function directly and :func:`option_probs` is only its | |
| exponentiated, detached presentation form. Keeping selection, validation, | |
| dtype conversion and temperature here prevents a train/serve readout fork. | |
| """ | |
| import torch | |
| if not math.isfinite(temperature) or temperature <= 0: | |
| raise ValueError("temperature must be finite and greater than zero") | |
| if not letter_token_ids: | |
| raise ValueError("letter_token_ids must not be empty") | |
| logits = torch.as_tensor(logits_last) | |
| if logits.ndim != 1: | |
| raise ValueError("logits_last must be a one-dimensional vocabulary vector") | |
| if any(isinstance(index, bool) or not isinstance(index, int) or not 0 <= index < logits.numel() | |
| for index in letter_token_ids): | |
| raise ValueError("letter_token_ids must be valid vocabulary indices") | |
| try: | |
| selected = logits[list(letter_token_ids)].to(dtype=torch.float32) / temperature | |
| except (IndexError, TypeError) as error: | |
| raise ValueError("letter_token_ids must be valid vocabulary indices") from error | |
| log_probabilities = torch.log_softmax(selected, dim=-1) | |
| if validate_finite and not bool(torch.isfinite(log_probabilities).all()): | |
| raise ValueError("option probabilities are non-finite") | |
| return log_probabilities | |
| def option_probs(logits_last, letter_token_ids: list[int], temperature: float = 1.0) -> list[float]: | |
| """FP32 softmax over only the displayed option-letter logits.""" | |
| values = option_log_probs(logits_last, letter_token_ids, temperature, validate_finite=False).exp().detach().cpu().tolist() | |
| if not all(math.isfinite(value) for value in values): | |
| raise ValueError("option probabilities are non-finite") | |
| return values | |