diff --git a/README.md b/README.md index a9ad53f..cc4f80d 100644 --- a/README.md +++ b/README.md @@ -90,6 +90,32 @@ ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \ --overwrite-output ``` +For QAT, `--quantization-aware` replaces the model's linear and embedding forwards with +straight-through fake quantization matching the shipping QDQ graph: per-tensor uint8 +embeddings, per-channel symmetric int8 linear weights, and per-tensor uint8 activations. +Parameter names remain unchanged, so the selected checkpoint reopens as an ordinary +Transformers model and uses the same `export.py` path. Keep the incoming checkpoint as +epoch zero and select QAT only on validation. Training logs progress every 50 batches by +default (`--progress-steps 0` disables it), so a long CPU run remains observable: + +```bash +ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \ + --model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \ + --epochs 2 --learning-rate 1e-6 --warmup-ratio 0 \ + --early-stopping-patience 1 --boundary-weight 2 --quantization-aware \ + --output-dir ml/purpose-classifier/outputs/purpose-lite-v1-qat1 \ + --overwrite-output +``` + +On dataset v1, that validation-selected run produced a 23,148,500-byte int8 graph at +94.88% frozen accuracy (889/937), 94.46% scored-hard accuracy, and 98.19% scorable +PyTorch↔ONNX agreement. It is the current quantized candidate, but remains two correct +predictions below the 95% gate. A subsequent validation-selected `5e-7` epoch improved +int8 validation accuracy from 93.31% to 93.71% but regressed frozen accuracy to 94.34%; +it is rejected. Do not continue optimizer-only QAT sweeps on this split. The next model +iteration should incorporate reviewed boundary data and be selected on a revised +validation/frozen dataset version. + For a wiring smoke test, use a small deterministic prefix: ```bash diff --git a/tests/test_train.py b/tests/test_train.py index 1f6958f..2eb0e20 100644 --- a/tests/test_train.py +++ b/tests/test_train.py @@ -12,6 +12,32 @@ import train class MetricsTests(unittest.TestCase): + def test_qat_replacements_keep_checkpoint_keys_and_gradients(self): + model = torch.nn.Sequential( + torch.nn.Embedding(16, 8), + torch.nn.Flatten(), + torch.nn.Linear(16, 3), + ) + keys = set(model.state_dict()) + counts = train.enable_quantization_aware_training(torch, model) + self.assertEqual({"linear": 1, "embedding": 1}, counts) + self.assertEqual(keys, set(model.state_dict())) + output = model(torch.tensor([[1, 2]], dtype=torch.long)) + output.sum().backward() + self.assertIsNotNone(model[0].weight.grad) + self.assertIsNotNone(model[2].weight.grad) + + def test_qat_affine_ranges_include_zero(self): + model = torch.nn.Sequential(torch.nn.Linear(2, 2, bias=False)) + with torch.no_grad(): + model[0].weight.copy_(torch.eye(2)) + train.enable_quantization_aware_training(torch, model) + output = model(torch.tensor([[1.0, 2.0]])) + self.assertTrue( + torch.allclose(output, torch.tensor([[1.0, 2.0]]), atol=0.02), + output, + ) + def test_boundary_training_weight_is_opt_in(self): self.assertEqual( 2.0, diff --git a/train.py b/train.py index 81e8b00..c9564e7 100644 --- a/train.py +++ b/train.py @@ -37,6 +37,123 @@ def training_weight(record: dict[str, Any], boundary_weight: float) -> float: return boundary_weight if record.get("slice") == "boundary" else 1.0 +def enable_quantization_aware_training(torch: Any, model: Any) -> dict[str, int]: + """Mirror the export graph's int8 policy with straight-through fake quantization. + + ONNX Runtime emits per-tensor uint8 embedding weights, per-channel symmetric int8 + linear weights, and per-tensor uint8 activations. The replacement modules keep the + original parameter names, so the selected checkpoint loads as an ordinary Transformers + model for export after fake-quantization-aware fine-tuning. + """ + + functional = torch.nn.functional + + def affine_parameters(value: Any) -> tuple[float, int]: + detached = value.detach().float() + # ONNX Runtime extends affine calibration ranges to include exact zero. + minimum = min(0.0, float(detached.amin().item())) + maximum = max(0.0, float(detached.amax().item())) + scale = max((maximum - minimum) / 255.0, torch.finfo(torch.float32).eps) + zero_point = max(0, min(255, round(-minimum / scale))) + return scale, zero_point + + def fake_quantize_activation(value: Any) -> Any: + scale, zero_point = affine_parameters(value) + return torch.fake_quantize_per_tensor_affine( + value, + scale, + zero_point, + 0, + 255, + ) + + def fake_quantize_linear_weight(weight: Any) -> Any: + detached = weight.detach().float() + scales = detached.abs().amax(dim=1).div(127.0).clamp_min( + torch.finfo(torch.float32).eps + ) + zero_points = torch.zeros_like(scales, dtype=torch.int32) + return torch.fake_quantize_per_channel_affine( + weight, + scales, + zero_points, + 0, + -127, + 127, + ) + + class QATLinear(torch.nn.Linear): + def forward(self, value: Any) -> Any: + result = functional.linear( + fake_quantize_activation(value), + fake_quantize_linear_weight(self.weight), + self.bias, + ) + return fake_quantize_activation(result) + + class QATEmbedding(torch.nn.Embedding): + def forward(self, indexes: Any) -> Any: + embedded = functional.embedding( + indexes, + self.weight, + self.padding_idx, + self.max_norm, + self.norm_type, + self.scale_grad_by_freq, + self.sparse, + ) + # Quantize only the selected rows using the full table's scale. This is + # numerically equivalent to dequantizing the whole table before Gather but + # avoids materializing a 30k x 384 fake-quantized embedding every batch. + weight_scale, weight_zero_point = affine_parameters(self.weight) + embedded = torch.fake_quantize_per_tensor_affine( + embedded, + weight_scale, + weight_zero_point, + 0, + 255, + ) + return fake_quantize_activation(embedded) + + counts = {"linear": 0, "embedding": 0} + + def replace(parent: Any) -> None: + for name, child in list(parent.named_children()): + replacement = None + if isinstance(child, torch.nn.Linear): + replacement = QATLinear( + child.in_features, + child.out_features, + bias=child.bias is not None, + device=child.weight.device, + dtype=child.weight.dtype, + ) + counts["linear"] += 1 + elif isinstance(child, torch.nn.Embedding): + replacement = QATEmbedding( + child.num_embeddings, + child.embedding_dim, + padding_idx=child.padding_idx, + max_norm=child.max_norm, + norm_type=child.norm_type, + scale_grad_by_freq=child.scale_grad_by_freq, + sparse=child.sparse, + device=child.weight.device, + dtype=child.weight.dtype, + ) + counts["embedding"] += 1 + if replacement is not None: + replacement.weight = child.weight + if isinstance(child, torch.nn.Linear): + replacement.bias = child.bias + setattr(parent, name, replacement) + else: + replace(child) + + replace(model) + return counts + + def encode_fixed_shape( tokenizer: Any, texts: Sequence[str], @@ -379,6 +496,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "headTokens": HEAD_TOKENS, "tailTokens": TAIL_TOKENS, } + config.purpose_classifier_quantization_aware_training = bool( + args.quantization_aware + ) + qat_modules = {"linear": 0, "embedding": 0} + if args.quantization_aware: + qat_modules = enable_quantization_aware_training(torch, model) model.to(device) class PromptDataset(Dataset): @@ -472,6 +595,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: started = time.perf_counter() for epoch in range(1, args.epochs + 1): + epoch_started = time.perf_counter() model.train() optimizer.zero_grad(set_to_none=True) running_loss = 0.0 @@ -498,6 +622,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]: optimizer.step() scheduler.step() optimizer.zero_grad(set_to_none=True) + if args.progress_steps and ( + step % args.progress_steps == 0 or step == len(train_loader) + ): + print( + f"epoch {epoch} step {step}/{len(train_loader)} " + f"mean_loss={running_loss / step:.4f} " + f"elapsed={time.perf_counter() - epoch_started:.1f}s", + flush=True, + ) logits, labels = _evaluate(torch, model, validation_loader, device) predictions = logits.argmax(dim=-1).tolist() @@ -586,6 +719,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "trainingSeconds": time.perf_counter() - started, "trainRecords": len(train_records), "boundaryTrainingWeight": args.boundary_weight, + "quantizationAwareTraining": args.quantization_aware, + "quantizationAwareModules": qat_modules, "validationRecords": len(validation_records), "scoredValidationRecords": int(validation_scorable.sum().item()), "vagueAbstentionValidationRecords": int( @@ -637,9 +772,11 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--warmup-ratio", type=float, default=0.1) parser.add_argument("--max-grad-norm", type=float, default=1.0) parser.add_argument("--workers", type=int, default=0) + parser.add_argument("--progress-steps", type=int, default=50) parser.add_argument("--early-stopping-patience", type=int, default=2) parser.add_argument("--minimum-improvement", type=float, default=0.0005) parser.add_argument("--boundary-weight", type=float, default=1.0) + parser.add_argument("--quantization-aware", action="store_true") parser.add_argument("--high-precision", type=float, default=0.98) parser.add_argument("--accepted-precision", type=float, default=0.95) parser.add_argument("--max-train-records", type=int) @@ -668,6 +805,8 @@ def main(argv: Sequence[str] | None = None) -> int: parser.error("--warmup-ratio must be in [0, 1)") if args.minimum_improvement < 0.0: parser.error("--minimum-improvement must be non-negative") + if args.progress_steps < 0: + parser.error("--progress-steps must be non-negative") if args.boundary_weight <= 0.0: parser.error("--boundary-weight must be positive") if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0: