Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user