File size: 6,704 Bytes
ad895de | 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 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | import os
# Configure workspace-local cache defaults to prevent permission errors on hosts
if "HF_HOME" not in os.environ:
os.environ["HF_HOME"] = os.path.expanduser("~/Nikola/.cache/huggingface")
if "HF_DATASETS_CACHE" not in os.environ:
os.environ["HF_DATASETS_CACHE"] = os.path.expanduser("~/Nikola/.cache/huggingface/datasets")
import sys
import random
import yaml
from pathlib import Path
import torch
from datasets import load_dataset, Dataset
from transformers import AutoModel, AutoTokenizer
from llmcompressor import oneshot
from llmcompressor.modifiers.quantization import QuantizationModifier
# Model ID configuration
MODEL_ID = "/home/olegk/Nikola/models/embedding/Qwen3-Embedding-4B"
SAVE_DIR = "/home/olegk/Nikola/models/embedding/Qwen3-Embedding-4B-NVFP4"
print(f"✨ Loading model: {MODEL_ID}")
model = AutoModel.from_pretrained(MODEL_ID, torch_dtype="auto")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
# Select number of samples.
NUM_CALIBRATION_SAMPLES = 512
MAX_SEQUENCE_LENGTH = 512
samples_per_dataset = NUM_CALIBRATION_SAMPLES // 2
# Extract local notes for domain calibration
notes_texts = []
notes_dir = Path.home() / ".sakura" / "notes"
if notes_dir.exists():
print("✨ Found local Sakura notes! Extracting domain-specific calibration data...")
for category in ["persons", "places", "topics"]:
cat_dir = notes_dir / category
if cat_dir.exists():
for f in cat_dir.glob("*.yaml"):
if category == "persons" and f.stem in ["oleg", "sakura"]:
continue
try:
with open(f, "r", encoding="utf-8") as file:
data = yaml.safe_load(file) or {}
except Exception:
continue
# Format text parts based on category
if category == "persons":
name = data.get("name", f.stem.title())
desc = data.get("short_description", "")
details = " ".join(data.get("details", []))
notes_texts.append(f"Person Profile: {name} - {desc} {details}")
if desc: notes_texts.append(desc)
if details: notes_texts.append(details)
elif category == "places":
name = data.get("name", f.stem.title())
desc = data.get("short_description", "")
details = " ".join(data.get("details", []))
notes_texts.append(f"Place Profile: {name} - {desc} {details}")
if desc: notes_texts.append(desc)
if details: notes_texts.append(details)
elif category == "topics":
topic = data.get("topic", f.stem.replace("_", " ").title())
words = data.get("words", [])
suggestions = " ".join(data.get("suggestions", []))
notes_texts.append(f"Vocabulary Topic: {topic}")
if suggestions: notes_texts.append(suggestions)
for w in words:
word = w.get("word")
trans = w.get("translation")
notes = w.get("notes", "")
if word and trans:
notes_texts.append(f"{word} : {trans}")
notes_texts.append(word)
notes_texts.append(trans)
if notes:
notes_texts.append(notes)
print(f"✨ Extracted {len(notes_texts)} local note calibration segments.")
print("✨ Building balanced bilingual calibration dataset...")
# Load English calibration set (from ultrachat_200k)
ds_en = load_dataset(
"HuggingFaceH4/ultrachat_200k",
split=f"train_sft[:{samples_per_dataset}]",
)
en_texts = [example["messages"][0]["content"] for example in ds_en]
# Load Japanese calibration set (from llm-jp-instructions)
ds_ja = load_dataset(
"llm-jp/llm-jp-instructions",
split=f"train[:{samples_per_dataset}]",
)
ja_texts = [example["text"] for example in ds_ja]
# Combine and shuffle
calibration_texts = notes_texts + en_texts + ja_texts
random.seed(42)
random.shuffle(calibration_texts)
# Slice to target sample count
calibration_texts = calibration_texts[:NUM_CALIBRATION_SAMPLES]
# Convert to HF Dataset object
ds = Dataset.from_dict({"text": calibration_texts})
# Tokenize inputs
def tokenize(sample):
return tokenizer(
sample["text"],
padding=False,
max_length=MAX_SEQUENCE_LENGTH,
truncation=True,
add_special_tokens=False,
)
print("✨ Tokenizing dataset...")
ds = ds.map(tokenize, remove_columns=ds.column_names)
# Custom data collator to ensure everything goes to cuda
def data_collator(batch):
assert len(batch) == 1
return {key: torch.tensor(value).unsqueeze(0).to(model.device) for key, value in batch[0].items()}
# Configure the quantization algorithm and scheme.
# We target only the MLP layers, leaving self-attention projections in bfloat16.
# Using NVFP4 (W4A4) ensures that calibration runs and generates the activation scales vLLM needs. 🌸✨
recipe = QuantizationModifier(
targets=[
"re:.*mlp.gate_proj$",
"re:.*mlp.up_proj$",
"re:.*mlp.down_proj$",
],
scheme="NVFP4",
)
print("✨ Running oneshot calibration & quantization to NVFP4 format...")
oneshot(
model=model,
dataset=ds,
recipe=recipe,
max_seq_length=MAX_SEQUENCE_LENGTH,
num_calibration_samples=NUM_CALIBRATION_SAMPLES,
data_collator=data_collator,
)
print(f"✨ Saving compressed model to {SAVE_DIR}...")
model.save_pretrained(SAVE_DIR, save_compressed=True)
tokenizer.save_pretrained(SAVE_DIR)
# Post-save patching to add regex ignores for vLLM compatibility.
# vLLM fuses attention layers into qkv_proj, which doesn't match the leaf-level
# ignores generated by llmcompressor. Add a regex to ignore all attention subtrees.
import json
config_path = os.path.join(SAVE_DIR, "config.json")
if os.path.exists(config_path):
with open(config_path, "r") as f:
config = json.load(f)
if "quantization_config" in config:
qc = config["quantization_config"]
ignore = qc.setdefault("ignore", [])
for pattern in ["re:.*self_attn.*"]:
if pattern not in ignore:
ignore.insert(0, pattern)
with open(config_path, "w") as f:
json.dump(config, f, indent=2)
print("✨ Successfully patched quantization_config.ignore for vLLM compatibility!")
print("🎉 Done! Model successfully quantized to NVFP4.")
|