Zero-Shot Classification
Safetensors
PEFT
English
openjev
classification
decision-model
listwise
gemma4
research
Instructions to use bambamdevs/openjev-e4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use bambamdevs/openjev-e4b with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
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
|