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
+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: