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
|
--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:
|
For a wiring smoke test, use a small deterministic prefix:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|||||||
@@ -12,6 +12,32 @@ import train
|
|||||||
|
|
||||||
|
|
||||||
class MetricsTests(unittest.TestCase):
|
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):
|
def test_boundary_training_weight_is_opt_in(self):
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
2.0,
|
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
|
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(
|
def encode_fixed_shape(
|
||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
texts: Sequence[str],
|
texts: Sequence[str],
|
||||||
@@ -379,6 +496,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"headTokens": HEAD_TOKENS,
|
"headTokens": HEAD_TOKENS,
|
||||||
"tailTokens": TAIL_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)
|
model.to(device)
|
||||||
|
|
||||||
class PromptDataset(Dataset):
|
class PromptDataset(Dataset):
|
||||||
@@ -472,6 +595,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
|
|
||||||
started = time.perf_counter()
|
started = time.perf_counter()
|
||||||
for epoch in range(1, args.epochs + 1):
|
for epoch in range(1, args.epochs + 1):
|
||||||
|
epoch_started = time.perf_counter()
|
||||||
model.train()
|
model.train()
|
||||||
optimizer.zero_grad(set_to_none=True)
|
optimizer.zero_grad(set_to_none=True)
|
||||||
running_loss = 0.0
|
running_loss = 0.0
|
||||||
@@ -498,6 +622,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
optimizer.step()
|
optimizer.step()
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
optimizer.zero_grad(set_to_none=True)
|
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)
|
logits, labels = _evaluate(torch, model, validation_loader, device)
|
||||||
predictions = logits.argmax(dim=-1).tolist()
|
predictions = logits.argmax(dim=-1).tolist()
|
||||||
@@ -586,6 +719,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"trainingSeconds": time.perf_counter() - started,
|
"trainingSeconds": time.perf_counter() - started,
|
||||||
"trainRecords": len(train_records),
|
"trainRecords": len(train_records),
|
||||||
"boundaryTrainingWeight": args.boundary_weight,
|
"boundaryTrainingWeight": args.boundary_weight,
|
||||||
|
"quantizationAwareTraining": args.quantization_aware,
|
||||||
|
"quantizationAwareModules": qat_modules,
|
||||||
"validationRecords": len(validation_records),
|
"validationRecords": len(validation_records),
|
||||||
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
||||||
"vagueAbstentionValidationRecords": int(
|
"vagueAbstentionValidationRecords": int(
|
||||||
@@ -637,9 +772,11 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
parser.add_argument("--warmup-ratio", type=float, default=0.1)
|
parser.add_argument("--warmup-ratio", type=float, default=0.1)
|
||||||
parser.add_argument("--max-grad-norm", type=float, default=1.0)
|
parser.add_argument("--max-grad-norm", type=float, default=1.0)
|
||||||
parser.add_argument("--workers", type=int, default=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("--early-stopping-patience", type=int, default=2)
|
||||||
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
|
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
|
||||||
parser.add_argument("--boundary-weight", type=float, default=1.0)
|
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("--high-precision", type=float, default=0.98)
|
||||||
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
||||||
parser.add_argument("--max-train-records", type=int)
|
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)")
|
parser.error("--warmup-ratio must be in [0, 1)")
|
||||||
if args.minimum_improvement < 0.0:
|
if args.minimum_improvement < 0.0:
|
||||||
parser.error("--minimum-improvement must be non-negative")
|
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:
|
if args.boundary_weight <= 0.0:
|
||||||
parser.error("--boundary-weight must be positive")
|
parser.error("--boundary-weight must be positive")
|
||||||
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
|
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
|
||||||
|
|||||||
Reference in New Issue
Block a user