Piko-9b / scripts /quantize_awq.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
4.83 kB
#!/usr/bin/env python3
"""Produce an AWQ build of Piko-9b.
NOT RUN for this release. No AWQ artefact has been produced or validated, and
nothing in the documentation claims one works.
Two architecture-specific hazards this script guards against:
1. **Linear-attention parameters are not ordinary linear weights.** `A_log`,
`dt_bias` and `conv1d` govern a recurrent state; quantizing them like a
`q_proj` can destabilise generation in ways a perplexity check will not catch.
They are excluded by default.
2. **Text-only calibration silently degrades the vision tower.** AWQ calibrates
on a text corpus. Applied to the vision tower and merger, that can break image
handling while every text metric stays healthy. The tower is left in bf16 by
default.
python scripts/quantize_awq.py --model Dexy2/Piko-9b --output ./piko-9b-awq
Afterwards you MUST run:
python scripts/validate_quantized_model.py --quantized ./piko-9b-awq \
--reference Dexy2/Piko-9b --output reports/quantization_validation.json
"""
from __future__ import annotations
import argparse
import json
import sys
import time
from pathlib import Path
# Modules whose quantization risks the recurrent state or the vision path.
EXCLUDED_PATTERNS = [
"visual", # entire vision tower and merger
"linear_attn.A_log",
"linear_attn.dt_bias",
"linear_attn.conv1d",
"linear_attn.norm",
"lm_head",
]
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--bits", type=int, default=4, choices=[4])
parser.add_argument("--group-size", type=int, default=128)
parser.add_argument("--zero-point", action="store_true", default=True)
parser.add_argument(
"--calibration-samples",
type=int,
default=128,
help="More samples give a better scale estimate and a slower run.",
)
parser.add_argument(
"--quantize-vision",
action="store_true",
help="Quantize the vision tower too. Unvalidated; expect image regressions.",
)
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
excluded = list(EXCLUDED_PATTERNS)
if args.quantize_vision:
excluded.remove("visual")
print(
"WARNING: quantizing the vision tower with text calibration. "
"Validate the image path before publishing.",
file=sys.stderr,
)
config = {
"zero_point": args.zero_point,
"q_group_size": args.group_size,
"w_bit": args.bits,
"version": "GEMM",
"modules_to_not_convert": excluded,
}
print("AWQ configuration:")
print(json.dumps(config, indent=2))
print(f"\nmodel : {args.model}")
print(f"output : {args.output}")
print(f"samples: {args.calibration_samples}")
print("\nEstimated cost: 30-90 minutes and >= 24 GB VRAM for a 9.65B model.")
if args.dry_run:
print("\n--dry-run: nothing executed.")
return
try:
from awq import AutoAWQForCausalLM
except ImportError:
sys.exit(
"autoawq is not installed:\n"
" pip install autoawq\n"
"Note: autoawq support for the qwen3_5 hybrid architecture has NOT been "
"verified. If it does not recognise the model type, this path is a dead end "
"until upstream adds support."
)
from transformers import AutoTokenizer
print("\nLoading model for quantization...", flush=True)
began = time.time()
model = AutoAWQForCausalLM.from_pretrained(args.model, device_map="cpu")
tokenizer = AutoTokenizer.from_pretrained(args.model)
print("Calibrating...", flush=True)
model.quantize(tokenizer, quant_config=config, max_calib_samples=args.calibration_samples)
args.output.mkdir(parents=True, exist_ok=True)
model.save_quantized(str(args.output))
tokenizer.save_pretrained(str(args.output))
(args.output / "quantization_provenance.json").write_text(
json.dumps(
{
"method": "awq",
"source_model": args.model,
"config": config,
"calibration_samples": args.calibration_samples,
"vision_quantized": args.quantize_vision,
"elapsed_seconds": round(time.time() - began, 1),
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
"validated": False,
},
indent=2,
)
+ "\n",
encoding="utf-8",
)
print(f"\nWrote {args.output}")
print("NOT YET VALIDATED. Run scripts/validate_quantized_model.py before publishing.")
if __name__ == "__main__":
main()