Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 18:39:33 -07:00
parent e4403b570d
commit ddbff97191
3 changed files with 191 additions and 0 deletions
+26
View File
@@ -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
+26
View File
@@ -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,
+139
View File
@@ -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: