File size: 5,914 Bytes
0810902
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
172
173
#!/usr/bin/env python3
"""Produce a GPTQ build of Piko-9b.

NOT RUN for this release. No GPTQ artefact has been produced or validated, and
nothing in the documentation claims one works.

Same two architecture-specific hazards as the AWQ path:

1. `A_log`, `dt_bias` and `conv1d` drive the linear-attention recurrent state and
   are not ordinary linear weights. Excluded by default.
2. GPTQ calibrates on text. Applying it to the vision tower can break image
   handling while text metrics stay healthy. The tower stays in bf16 by default.

    python scripts/quantize_gptq.py --model Dexy2/Piko-9b --output ./piko-9b-gptq

Afterwards you MUST run:

    python scripts/validate_quantized_model.py --quantized ./piko-9b-gptq \
        --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

EXCLUDED_PATTERNS = [
    "visual",
    "linear_attn.A_log",
    "linear_attn.dt_bias",
    "linear_attn.conv1d",
    "linear_attn.norm",
    "lm_head",
]

CALIBRATION_PROMPTS = [
    "Explain how a hash table resolves collisions.",
    "Summarise the causes of the 1929 financial crash.",
    "Write a Python function that merges two sorted lists.",
    "Describe the water cycle in four sentences.",
    "What is the difference between TCP and UDP?",
    "Extract the total from an invoice and return it as JSON.",
    "Explain gradient clipping and when it helps.",
    "Rewrite this sentence in the passive voice: The cat chased the mouse.",
]


def build_calibration(tokenizer, samples: int, seq_len: int) -> list[dict]:
    """A small, self-contained calibration set: no dataset download, no licence question."""
    texts = []
    while len(texts) < samples:
        for prompt in CALIBRATION_PROMPTS:
            texts.append(prompt)
            if len(texts) >= samples:
                break
    return [
        tokenizer(text, return_tensors="pt", truncation=True, max_length=seq_len) for text in texts
    ]


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=[2, 3, 4, 8])
    parser.add_argument("--group-size", type=int, default=128)
    parser.add_argument("--damp-percent", type=float, default=0.01)
    parser.add_argument(
        "--desc-act", action="store_true", default=False, help="Better accuracy, slower inference."
    )
    parser.add_argument("--calibration-samples", type=int, default=128)
    parser.add_argument("--sequence-length", type=int, default=2048)
    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 = {
        "bits": args.bits,
        "group_size": args.group_size,
        "damp_percent": args.damp_percent,
        "desc_act": args.desc_act,
        "sym": True,
        "true_sequential": True,
        "modules_to_not_convert": excluded,
    }

    print("GPTQ configuration:")
    print(json.dumps(config, indent=2))
    print(f"\nmodel  : {args.model}")
    print(f"output : {args.output}")
    print(f"samples: {args.calibration_samples} @ {args.sequence_length} tokens")
    print("\nEstimated cost: 1-3 hours and >= 24 GB VRAM for a 9.65B model.")

    if args.dry_run:
        print("\n--dry-run: nothing executed.")
        return

    try:
        from gptqmodel import GPTQModel, QuantizeConfig
    except ImportError:
        sys.exit(
            "gptqmodel is not installed:\n"
            "    pip install gptqmodel\n"
            "Note: gptqmodel 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

    tokenizer = AutoTokenizer.from_pretrained(args.model)
    calibration = build_calibration(tokenizer, args.calibration_samples, args.sequence_length)

    quantize_config = QuantizeConfig(
        bits=args.bits,
        group_size=args.group_size,
        damp_percent=args.damp_percent,
        desc_act=args.desc_act,
    )

    print("\nLoading model...", flush=True)
    began = time.time()
    model = GPTQModel.load(args.model, quantize_config)

    print("Quantizing...", flush=True)
    model.quantize(calibration)

    args.output.mkdir(parents=True, exist_ok=True)
    model.save(str(args.output))
    tokenizer.save_pretrained(str(args.output))

    (args.output / "quantization_provenance.json").write_text(
        json.dumps(
            {
                "method": "gptq",
                "source_model": args.model,
                "config": config,
                "calibration_samples": args.calibration_samples,
                "calibration_source": "self-contained prompt list (no external dataset)",
                "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()