Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
235 changes: 235 additions & 0 deletions compress_all.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,235 @@
import argparse
import sys
from pathlib import Path

import torch
from compressed_tensors.offload import init_dist
from compressed_tensors.quantization import (
QuantizationArgs,
QuantizationScheme,
QuantizationStrategy,
QuantizationType,
)
from compressed_tensors.quantization.quant_args import FP8_E4M3_DATA
from transformers import AutoModelForCausalLM, AutoTokenizer

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

# ── Constants ────────────────────────────────────────────────────────

DEFAULT_MODELS = [
"meta-llama/Meta-Llama-3-8B-Instruct",
"meta-llama/Meta-Llama-3-70B-Instruct",
]

OUTPUT_BASE = Path("/tmp/fouroversix-sanity")

COMMON = dict(
num_bits=4,
type=QuantizationType.FLOAT,
strategy=QuantizationStrategy.TENSOR_GROUP,
symmetric=True,
dynamic=False,
group_size=16,
scale_dtype=FP8_E4M3_DATA.dtype,
zp_dtype=FP8_E4M3_DATA.dtype,
)

# ── Observer configs ─────────────────────────────────────────────────

CONFIGS = {
"fouroversix": QuantizationArgs(
**COMMON,
observer="fouroversix",
),
"mse-1x-1.5x": QuantizationArgs(
**COMMON,
observer="memoryless_mse",
observer_kwargs={
"expand": 1.5,
"grid": 3,
"maxshrink": 0.67,
"norm": 2.0,
"patience": 100000,
},
),
"nvfp4_expanded_mse": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
),
"expand-3.4": QuantizationArgs(
**COMMON,
observer="memoryless_mse",
observer_kwargs={
"expand": 3.4,
"maxshrink": round(1 - 0.8 / 3.4, 4),
"grid": 200.0,
"patience": 200,
},
),
"expanded-norm1.8": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
observer_kwargs={"norm": 1.8},
),
"expanded-norm2.0": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
observer_kwargs={"norm": 2.0},
),
"expanded-norm2.2": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
observer_kwargs={"norm": 2.2},
),
"expanded-norm2.4": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
observer_kwargs={"norm": 2.4},
),
"default-mse": QuantizationArgs(
**COMMON,
observer="memoryless_mse",
),
"minmax": QuantizationArgs(
**COMMON,
observer="memoryless_minmax",
),
"expanded-gs-prior": QuantizationArgs(
**COMMON,
observer="nvfp4_expanded_mse",
observer_kwargs={"use_global_scale_prior": True},
),
}

CONFIG_DESCRIPTIONS = {
"fouroversix": "FourOverSix: per-block M=6/M=4 adaptive scaling, gs_max=256",
"mse-1x-1.5x": "MSE 1x+1.5x: 2-point search matching FourOverSix search space",
"nvfp4_expanded_mse": "NVFP4 Expanded MSE: 1.8x→0.8x range, 112 steps (default norm=2.4)",
"expand-3.4": "Original ablation: expand=3.4, maxshrink=0.7647, grid=200, patience=200",
"expanded-norm1.8": "NVFP4 Expanded MSE: norm=1.8",
"expanded-norm2.0": "NVFP4 Expanded MSE: norm=2.0",
"expanded-norm2.2": "NVFP4 Expanded MSE: norm=2.2",
"expanded-norm2.4": "NVFP4 Expanded MSE: norm=2.4 (explicit)",
"default-mse": "Default MSE: expand=1.0, maxshrink=0.20, grid=100, norm=2.4",
"minmax": "MinMax: no MSE search, simple min/max scaling",
"expanded-gs-prior": "NVFP4 Expanded MSE: with global_scale prior in search",
}


# ── Helpers ──────────────────────────────────────────────────────────


def model_short_name(model_id: str) -> str:
return model_id.rstrip("/").split("/")[-1]


def compress(
model_id: str,
config_key: str,
output_base: Path,
force: bool = False,
):
name = model_short_name(model_id)
out = output_base / f"{name}-{config_key}"

if out.exists() and not force:
print(f" SKIP (exists): {out}")
return out

print(f"\n{'='*70}")
print(f" {out}")
print(f" model: {model_id}")
print(f" config: {config_key} — {CONFIG_DESCRIPTIONS.get(config_key, '')}")
print(f"{'='*70}\n")

tokenizer = AutoTokenizer.from_pretrained(model_id)
with load_context():
model = AutoModelForCausalLM.from_pretrained(
model_id, device_map="auto_offload"
)

recipe = QuantizationModifier(
config_groups={
"group_0": QuantizationScheme(
targets=["Linear"],
weights=CONFIGS[config_key],
)
},
ignore=["lm_head"],
)

oneshot(
model=model,
recipe=recipe,
output_dir=str(out),
)
tokenizer.save_pretrained(out)
del model
torch.cuda.empty_cache()
print(f" DONE: {out}\n")
return out


# ── Main ─────────────────────────────────────────────────────────────


def main():
parser = argparse.ArgumentParser(
description="Compress models with NVFP4 observer configs"
)
parser.add_argument(
"--models", nargs="+", default=DEFAULT_MODELS,
help="HuggingFace model IDs to compress",
)
parser.add_argument(
"--configs", nargs="+", default=list(CONFIGS.keys()),
help="Observer configs to run",
)
parser.add_argument(
"--output-dir", type=str, default=str(OUTPUT_BASE),
help="Base directory for compressed models",
)
parser.add_argument(
"--force", action="store_true",
help="Recompress even if output exists",
)
parser.add_argument(
"--list", action="store_true",
help="List configs and exit",
)
args = parser.parse_args()

if args.list:
print("Available configs:")
for key, desc in CONFIG_DESCRIPTIONS.items():
print(f" {key:<25s} {desc}")
print(f"\nDefault models: {', '.join(DEFAULT_MODELS)}")
sys.exit(0)

for key in args.configs:
if key not in CONFIGS:
print(f"Unknown config: {key}")
print(f"Available: {', '.join(CONFIGS.keys())}")
sys.exit(1)

output_base = Path(args.output_dir)
output_base.mkdir(parents=True, exist_ok=True)

init_dist()

total = len(args.models) * len(args.configs)
idx = 0
for model_id in args.models:
for config_key in args.configs:
idx += 1
print(f"\n[{idx}/{total}] {model_short_name(model_id)} / {config_key}")
compress(model_id, config_key, output_base, args.force)

torch.distributed.destroy_process_group()


if __name__ == "__main__":
main()
158 changes: 158 additions & 0 deletions eval_all.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
import argparse
import sys
from pathlib import Path

import torch
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer

# ── Constants ────────────────────────────────────────────────────────

DEFAULT_MODELS = [
"meta-llama/Meta-Llama-3-8B-Instruct",
"meta-llama/Meta-Llama-3-70B-Instruct",
]

DEFAULT_CONFIGS = [
"fouroversix",
"mse-1x-1.5x",
"nvfp4_expanded_mse",
"expand-3.4",
]

COMPRESSED_DIR = Path("/tmp/fouroversix-sanity")

CONFIG_LABELS = {
"fouroversix": "FourOverSix",
"mse-1x-1.5x": "MSE 1x+1.5x",
"nvfp4_expanded_mse": "nvfp4_expanded_mse",
"expand-3.4": "expand=3.4 (original ablation)",
"default-mse": "Default MSE",
"minmax": "MinMax",
"expanded-gs-prior": "Expanded MSE + gs prior",
}


# ── Evaluation ───────────────────────────────────────────────────────


def evaluate_ppl(model, tokenizer, max_length=2048, stride=512):
testdata = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split="test")
text = "\n\n".join(testdata["text"])
encodings = tokenizer(text, return_tensors="pt")
seq_len = encodings.input_ids.size(1)

nlls = []
prev_end_loc = 0
for begin_loc in range(0, seq_len, stride):
end_loc = min(begin_loc + max_length, seq_len)
trg_len = end_loc - prev_end_loc
device = next(model.parameters()).device
input_ids = encodings.input_ids[:, begin_loc:end_loc].to(device)
target_ids = input_ids.clone()
target_ids[:, :-trg_len] = -100
with torch.no_grad():
outputs = model(input_ids, labels=target_ids)
nlls.append(outputs.loss)
prev_end_loc = end_loc
if end_loc == seq_len:
break

return torch.exp(torch.stack(nlls).mean()).item()


def load_and_eval(path: str, label: str) -> float:
print(f" Loading {label} from {path}...", flush=True)
tokenizer = AutoTokenizer.from_pretrained(path)
model = AutoModelForCausalLM.from_pretrained(
path, device_map="auto", torch_dtype="auto"
)
model.eval()
ppl = evaluate_ppl(model, tokenizer)
print(f" {label}: {ppl:.3f}")
del model
torch.cuda.empty_cache()
return ppl


# ── Main ─────────────────────────────────────────────────────────────


def model_short_name(model_id: str) -> str:
return model_id.rstrip("/").split("/")[-1]


def main():
parser = argparse.ArgumentParser(
description="PPL evaluation for NVFP4 observer comparison"
)
parser.add_argument(
"--models", nargs="+", default=DEFAULT_MODELS,
help="HuggingFace model IDs (baselines)",
)
parser.add_argument(
"--configs", nargs="+", default=DEFAULT_CONFIGS,
help="Observer config names to evaluate",
)
parser.add_argument(
"--compressed-dir", type=str, default=str(COMPRESSED_DIR),
help="Directory containing compressed models",
)
parser.add_argument(
"--no-baseline", action="store_true",
help="Skip unquantized baseline evaluation",
)
args = parser.parse_args()

compressed_dir = Path(args.compressed_dir)
all_results = {}

for model_id in args.models:
name = model_short_name(model_id)
print(f"\n{'='*60}")
print(f" {name}")
print(f"{'='*60}")

results = []
baseline_ppl = None

if not args.no_baseline:
baseline_ppl = load_and_eval(model_id, "Unquantized")
results.append(("Unquantized", baseline_ppl))

for config_key in args.configs:
path = compressed_dir / f"{name}-{config_key}"
label = CONFIG_LABELS.get(config_key, config_key)
if not path.exists():
print(f" SKIP (not found): {path}")
results.append((label, None))
continue
ppl = load_and_eval(str(path), label)
results.append((label, ppl))

all_results[name] = (baseline_ppl, results)

# ── Print markdown tables ────────────────────────────────────
print(f"\n\n{'='*60}")
print("Results (markdown)")
print(f"{'='*60}\n")

for name, (baseline_ppl, results) in all_results.items():
print(f"### {name}\n")
print("| Config | word_perplexity | delta vs unquantized |")
print("|---|---|---|")
for label, ppl in results:
if ppl is None:
print(f"| {label} | — | — |")
elif label == "Unquantized":
print(f"| {label} | **{ppl:.3f}** | — |")
elif baseline_ppl is not None:
delta = ppl - baseline_ppl
print(f"| {label} | {ppl:.3f} | +{delta:.3f} |")
else:
print(f"| {label} | {ppl:.3f} | — |")
print()


if __name__ == "__main__":
main()
1 change: 1 addition & 0 deletions src/llmcompressor/observers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,3 +15,4 @@
from .min_max import *
from .mse import *
from .imatrix import *
from .fouroversix import *
Loading
Loading