YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

# FP8 static per-tensor quantization for google/gemma-4-26B-A4B-it
#
# Weights: fp8 per-tensor static
# Activations: fp8 per-tensor static (calibrated via static_minmax observer)

import os

os.environ["HF_HOME"] = os.path.expanduser("~/hf_hub")
os.environ["HF_DATASETS_CACHE"] = "/tmp/hf_datasets_cache"

import torch
import datasets
from datasets import load_dataset

datasets.disable_caching()
from transformers import AutoProcessor, Gemma4ForConditionalGeneration

from llmcompressor import oneshot
from llmcompressor.modifiers.quantization import QuantizationModifier
from llmcompressor.utils import load_context

MODEL_ID = "google/gemma-4-26B-A4B-it"

# Load model.
with load_context(Gemma4ForConditionalGeneration):
    model = Gemma4ForConditionalGeneration.from_pretrained(MODEL_ID)
processor = AutoProcessor.from_pretrained(MODEL_ID)

# MoE expert handling is applied automatically by the pipeline.

# Configure FP8 static per-tensor quantization.
# Weights are quantized per-tensor, activations use static scales
# computed from calibration data via the static_minmax observer.
recipe = QuantizationModifier(
    targets="Linear",
    scheme="FP8",
    ignore=[
        "lm_head",
        "re:.*embed.*",
        "re:.*router",
        "re:.*vision_tower.*",
    ],
)

# Calibration dataset (needed for static activation scales).
DATASET_ID = "neuralmagic/calibration"
NUM_CALIBRATION_SAMPLES = 512
MAX_SEQUENCE_LENGTH = 8192

ds = load_dataset(DATASET_ID, name="LLM", split=f"train[:{NUM_CALIBRATION_SAMPLES}]")


def preprocess_function(example):
    messages = []
    for message in example["messages"]:
        messages.append(
            {
                "role": message["role"],
                "content": [{"type": "text", "text": message["content"]}],
            }
        )

    return processor.apply_chat_template(
        messages,
        return_tensors="pt",
        padding=False,
        truncation=True,
        max_length=MAX_SEQUENCE_LENGTH,
        tokenize=True,
        add_special_tokens=False,
        return_dict=True,
        add_generation_prompt=False,
    )


ds = ds.map(preprocess_function, batched=False, remove_columns=ds.column_names)


def data_collator(batch):
    assert len(batch) == 1
    return {
        key: (
            torch.tensor(value)
            if key != "pixel_values"
            else torch.tensor(value, dtype=torch.bfloat16).squeeze(0)
        )
        for key, value in batch[0].items()
    }


# Apply quantization.
oneshot(
    model=model,
    recipe=recipe,
    dataset=ds,
    max_seq_length=MAX_SEQUENCE_LENGTH,
    num_calibration_samples=NUM_CALIBRATION_SAMPLES,
    data_collator=data_collator,
)

# Save to disk in compressed-tensors format.
SAVE_DIR = MODEL_ID.rstrip("/").split("/")[-1] + "-FP8-Static"
model.save_pretrained(SAVE_DIR, save_compressed=True)
processor.save_pretrained(SAVE_DIR)
Downloads last month
258
Safetensors
Model size
26B params
Tensor type
BF16
·
F8_E4M3
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support