"""Experimental image and audio inputs. The adapter, backbone delta and decision head were trained on text only. This path inserts Gemma 4's own image or audio soft tokens at the start of the state and reads the same decision positions as text. It has passed small smoke tests only; see docs/MULTIMODAL_EXPERIMENTAL.md. """ from __future__ import annotations import json import wave from pathlib import Path from types import SimpleNamespace import numpy as np from huggingface_hub import hf_hub_download from PIL import Image SAMPLING_RATE = 16000 ENCODER_MODULES = ["vision_tower", "audio_tower", "embed_vision", "embed_audio"] def quantization_skip_modules(model_id: str, revision: str) -> list[str]: """Transformers' default NF4 skip list plus the vision/audio encoders. Setting llm_int8_skip_modules replaces the default list, so it is rebuilt from a weightless copy of the model; the text layers stay quantized exactly as in the text-only path. """ import torch from transformers import AutoConfig, AutoModelForMultimodalLM from transformers.quantizers.base import get_keys_to_not_convert config = AutoConfig.from_pretrained(model_id, revision=revision) with torch.device("meta"): skeleton = AutoModelForMultimodalLM.from_config(config) # Skip patterns match from the start of the full module name, e.g. "model.audio_tower". encoders = {name for name, _ in skeleton.named_modules() if name.split(".")[-1] in ENCODER_MODULES} if len(encoders) != len(ENCODER_MODULES): raise RuntimeError(f"Expected encoder modules {ENCODER_MODULES}, found {sorted(encoders)}") return sorted(set(get_keys_to_not_convert(skeleton)) | encoders) def load_image(image) -> Image.Image: if isinstance(image, Image.Image): return image.convert("RGB") with Image.open(Path(image)) as im: return im.convert("RGB") def load_audio(audio, sampling_rate: int | None = None) -> np.ndarray: """Return mono float32 samples at 16 kHz from a PCM .wav path or an array.""" if isinstance(audio, (str, Path)): with wave.open(str(audio), "rb") as w: width, channels, rate = w.getsampwidth(), w.getnchannels(), w.getframerate() raw = w.readframes(w.getnframes()) if width == 1: samples = (np.frombuffer(raw, dtype=np.uint8).astype(np.float32) - 128.0) / 128.0 elif width == 2: samples = np.frombuffer(raw, dtype="