File size: 6,093 Bytes
03223d7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""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="<i2").astype(np.float32) / 32768.0
        elif width == 4:
            samples = np.frombuffer(raw, dtype="<i4").astype(np.float32) / 2147483648.0
        else:
            raise ValueError(f"Unsupported .wav sample width: {width} bytes")
        samples = samples.reshape(-1, channels).mean(axis=1)
    else:
        if sampling_rate is None:
            raise ValueError("Pass sampling_rate with an audio array")
        samples, rate = np.asarray(audio, dtype=np.float32), int(sampling_rate)
        if samples.ndim == 2:
            samples = samples.mean(axis=0 if samples.shape[0] < samples.shape[1] else 1)
    if rate != SAMPLING_RATE:
        n = int(round(len(samples) * SAMPLING_RATE / rate))
        samples = np.interp(np.linspace(0, len(samples) - 1, n), np.arange(len(samples)), samples)
    return samples.astype(np.float32)


class MediaEncoder:
    """Turns an image and/or audio clip into Gemma 4 soft-token ids and model inputs."""

    def __init__(self, model_id: str, revision: str, tokenizer):
        from transformers.models.gemma4.feature_extraction_gemma4 import Gemma4AudioFeatureExtractor
        from transformers.models.gemma4.image_processing_pil_gemma4 import Gemma4ImageProcessorPil
        from transformers.models.gemma4.processing_gemma4 import Gemma4Processor

        cfg = json.loads(Path(hf_hub_download(model_id, "processor_config.json", revision=revision)).read_text(encoding="utf-8"))
        self.image_processor = Gemma4ImageProcessorPil(
            **{k: v for k, v in cfg["image_processor"].items() if k != "image_processor_type"})
        self.feature_extractor = Gemma4AudioFeatureExtractor(
            **{k: v for k, v in cfg["feature_extractor"].items() if k != "feature_extractor_type"})
        # Reuse Gemma's exact audio-length arithmetic so placeholders match encoder output.
        self._audio_cfg = SimpleNamespace(audio_seq_length=int(cfg.get("audio_seq_length", 750)))
        self._count_audio_tokens = Gemma4Processor._compute_audio_num_tokens
        ids = tokenizer.convert_tokens_to_ids
        self.image_token_id = ids(tokenizer.image_token)
        self.audio_token_id = ids(tokenizer.audio_token)
        self._image_wrap = (ids(tokenizer.boi_token), ids(tokenizer.eoi_token))
        self._audio_wrap = (ids(tokenizer.boa_token), ids(tokenizer.eoa_token))
        self._newline = tokenizer("\n", add_special_tokens=False).input_ids

    def encode(self, image=None, audio=None, sampling_rate: int | None = None):
        """Return (token ids to insert into the state, extra model inputs)."""
        ids, inputs = [], {}
        if image is not None:
            out = self.image_processor(images=[load_image(image)], return_tensors="pt")
            n = int(out["num_soft_tokens_per_image"][0])
            ids += [self._image_wrap[0]] + [self.image_token_id] * n + [self._image_wrap[1]] + self._newline
            inputs["pixel_values"] = out["pixel_values"]
            inputs["image_position_ids"] = out["image_position_ids"]
        if audio is not None:
            samples = load_audio(audio, sampling_rate)
            n = self._count_audio_tokens(self._audio_cfg, samples, SAMPLING_RATE)
            if n <= 0:
                raise ValueError("Audio clip is too short")
            out = self.feature_extractor([samples], sampling_rate=SAMPLING_RATE, return_tensors="pt")
            ids += [self._audio_wrap[0]] + [self.audio_token_id] * n + [self._audio_wrap[1]] + self._newline
            inputs["input_features"] = out["input_features"]
            inputs["input_features_mask"] = out["input_features_mask"]
        return ids, inputs