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
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support