openjev-e4b / openjev /multimodal.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
6.09 kB
"""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