Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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