Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -331,7 +331,8 @@ To continue a completed run without discarding its trained task heads, pass its
|
|||||||
Continuation restores the backbone and all four heads strictly, then starts a fresh
|
Continuation restores the backbone and all four heads strictly, then starts a fresh
|
||||||
optimizer and learning-rate schedule; `--model` remains reserved for an untrained local
|
optimizer and learning-rate schedule; `--model` remains reserved for an untrained local
|
||||||
upstream checkpoint. The first base run was still improving when its three-epoch schedule
|
upstream checkpoint. The first base run was still improving when its three-epoch schedule
|
||||||
ended, so continue its selected checkpoint conservatively before changing architecture:
|
ended, so its selected checkpoint was continued conservatively before changing
|
||||||
|
architecture:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
ml/purpose-classifier/venv/bin/python -u \
|
ml/purpose-classifier/venv/bin/python -u \
|
||||||
@@ -350,6 +351,58 @@ ml/purpose-classifier/venv/bin/python -u \
|
|||||||
--overwrite-output
|
--overwrite-output
|
||||||
```
|
```
|
||||||
|
|
||||||
|
That continuation reached 80.31% primary and 78.28% hard-slice validation accuracy;
|
||||||
|
calibrated mixed F1 reached 67.12%. It remained far below purpose-lite, while primary
|
||||||
|
training loss and validation accuracy were still improving. Do not chain another plain
|
||||||
|
continuation. The bounded next experiment distills the mature purpose-lite boundary
|
||||||
|
teacher into the continued deep checkpoint while retaining direct primary labels and all
|
||||||
|
three auxiliary losses.
|
||||||
|
|
||||||
|
Create a teacher cache bound to the history-augmented split:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ml/purpose-classifier/venv/bin/python -u \
|
||||||
|
ml/purpose-classifier/cache_teacher.py \
|
||||||
|
--dataset-dir \
|
||||||
|
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
|
||||||
|
--model \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
|
||||||
|
--output \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
|
||||||
|
--device mps \
|
||||||
|
--batch-size 16 \
|
||||||
|
--progress-steps 25 \
|
||||||
|
--overwrite-output
|
||||||
|
```
|
||||||
|
|
||||||
|
Then run two validation-selected distilled continuation epochs:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ml/purpose-classifier/venv/bin/python -u \
|
||||||
|
ml/purpose-classifier/train_deep_mlx.py \
|
||||||
|
--variant base \
|
||||||
|
--resume-from \
|
||||||
|
ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-cont-3e/model \
|
||||||
|
--dataset-dir \
|
||||||
|
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
|
||||||
|
--distillation-cache \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
|
||||||
|
--distillation-weight 0.5 \
|
||||||
|
--distillation-temperature 2 \
|
||||||
|
--epochs 2 \
|
||||||
|
--learning-rate 1e-5 \
|
||||||
|
--early-stopping-patience 1 \
|
||||||
|
--progress-steps 10 \
|
||||||
|
--output-dir \
|
||||||
|
ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-distilled \
|
||||||
|
--overwrite-output
|
||||||
|
```
|
||||||
|
|
||||||
|
The trainer records the resumed checkpoint as epoch zero before updating anything, so a
|
||||||
|
distillation regression cannot overwrite the 80.31% candidate. Teacher agreement is
|
||||||
|
reported for diagnosis but does not enter deep checkpoint selection; overall and hard
|
||||||
|
primary label accuracy remain the only selection inputs.
|
||||||
|
|
||||||
Do not launch the large rung yet. It is justified only after base is evaluated on the
|
Do not launch the large rung yet. It is justified only after base is evaluated on the
|
||||||
frozen set; large must beat base by at least two hard-slice points, while deep itself must
|
frozen set; large must beat base by at least two hard-slice points, while deep itself must
|
||||||
reach 97% scored overall and beat the shipping lite artifact by five hard-slice points.
|
reach 97% scored overall and beat the shipping lite artifact by five hard-slice points.
|
||||||
|
|||||||
@@ -274,6 +274,20 @@ def multitask_metrics(
|
|||||||
targets.secondary[targets.secondary_mask].tolist(),
|
targets.secondary[targets.secondary_mask].tolist(),
|
||||||
secondary_predictions[targets.secondary_mask].tolist(),
|
secondary_predictions[targets.secondary_mask].tolist(),
|
||||||
)
|
)
|
||||||
|
supported_labels = [
|
||||||
|
label
|
||||||
|
for label, row in zip(
|
||||||
|
secondary["confusionMatrix"]["labels"],
|
||||||
|
secondary["confusionMatrix"]["rows"],
|
||||||
|
)
|
||||||
|
if sum(row)
|
||||||
|
]
|
||||||
|
secondary["supportedLabels"] = supported_labels
|
||||||
|
secondary["supportedMacroRecall"] = float(
|
||||||
|
np.mean(
|
||||||
|
[secondary["perPurposeRecall"][label] for label in supported_labels]
|
||||||
|
)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
secondary = None
|
secondary = None
|
||||||
|
|
||||||
|
|||||||
@@ -121,6 +121,8 @@ class TargetAndMetricTests(unittest.TestCase):
|
|||||||
self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"])
|
self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"])
|
||||||
self.assertAlmostEqual(1 / 6, metrics["selectionScore"])
|
self.assertAlmostEqual(1 / 6, metrics["selectionScore"])
|
||||||
self.assertEqual(1, metrics["secondary"]["accuracy"])
|
self.assertEqual(1, metrics["secondary"]["accuracy"])
|
||||||
|
self.assertEqual(["writing"], metrics["secondary"]["supportedLabels"])
|
||||||
|
self.assertEqual(1, metrics["secondary"]["supportedMacroRecall"])
|
||||||
self.assertEqual(1, metrics["mixed"]["f1"])
|
self.assertEqual(1, metrics["mixed"]["f1"])
|
||||||
|
|
||||||
def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self):
|
def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self):
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from deep_model_mlx import (
|
|||||||
save_weights,
|
save_weights,
|
||||||
)
|
)
|
||||||
from purpose_data import DataError, LABELS
|
from purpose_data import DataError, LABELS
|
||||||
|
from train_deep_mlx import _distillation_loss
|
||||||
|
|
||||||
|
|
||||||
def tiny_config(*, checkpointing=False):
|
def tiny_config(*, checkpointing=False):
|
||||||
@@ -97,6 +98,23 @@ class DeepModelTests(unittest.TestCase):
|
|||||||
mx.eval(value, gradients)
|
mx.eval(value, gradients)
|
||||||
self.assertTrue(float(value.item()) > 0)
|
self.assertTrue(float(value.item()) > 0)
|
||||||
|
|
||||||
|
def test_distillation_loss_matches_teacher_and_backpropagates(self):
|
||||||
|
teacher = mx.array([[2.0, 0.0, -1.0]])
|
||||||
|
|
||||||
|
def loss(student):
|
||||||
|
return _distillation_loss(
|
||||||
|
mx,
|
||||||
|
student,
|
||||||
|
teacher,
|
||||||
|
temperature=2.0,
|
||||||
|
weights=mx.ones((1,)),
|
||||||
|
)
|
||||||
|
|
||||||
|
value, gradient = mx.value_and_grad(loss)(teacher)
|
||||||
|
mx.eval(value, gradient)
|
||||||
|
self.assertAlmostEqual(0.0, float(value.item()), places=6)
|
||||||
|
self.assertEqual(teacher.shape, gradient.shape)
|
||||||
|
|
||||||
def test_checkpoint_round_trip(self):
|
def test_checkpoint_round_trip(self):
|
||||||
model = ModernBertForPurposeClassification(tiny_config())
|
model = ModernBertForPurposeClassification(tiny_config())
|
||||||
with tempfile.TemporaryDirectory() as temp:
|
with tempfile.TemporaryDirectory() as temp:
|
||||||
|
|||||||
+160
-14
@@ -37,7 +37,7 @@ from train import (
|
|||||||
choose_confidence_thresholds,
|
choose_confidence_thresholds,
|
||||||
expected_calibration_error,
|
expected_calibration_error,
|
||||||
)
|
)
|
||||||
from train_mlx import _configure_mlx_device, _linear_schedule
|
from train_mlx import _configure_mlx_device, _linear_schedule, _teacher_cache
|
||||||
|
|
||||||
|
|
||||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||||
@@ -254,6 +254,39 @@ def _softmax(values: np.ndarray) -> np.ndarray:
|
|||||||
return exponentials / exponentials.sum(axis=-1, keepdims=True)
|
return exponentials / exponentials.sum(axis=-1, keepdims=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _distillation_loss(
|
||||||
|
mx: Any,
|
||||||
|
student_logits: Any,
|
||||||
|
teacher_logits: Any,
|
||||||
|
*,
|
||||||
|
temperature: float,
|
||||||
|
weights: Any,
|
||||||
|
) -> Any:
|
||||||
|
"""Return weighted teacher-to-student KL loss for one MLX batch."""
|
||||||
|
|
||||||
|
student_log_probabilities = (
|
||||||
|
student_logits / temperature
|
||||||
|
- mx.logsumexp(student_logits / temperature, axis=-1, keepdims=True)
|
||||||
|
)
|
||||||
|
teacher_probabilities = mx.softmax(
|
||||||
|
teacher_logits / temperature,
|
||||||
|
axis=-1,
|
||||||
|
)
|
||||||
|
teacher_log_probabilities = mx.log(
|
||||||
|
mx.maximum(teacher_probabilities, 1e-12)
|
||||||
|
)
|
||||||
|
per_record = (
|
||||||
|
mx.sum(
|
||||||
|
teacher_probabilities
|
||||||
|
* (teacher_log_probabilities - student_log_probabilities),
|
||||||
|
axis=-1,
|
||||||
|
)
|
||||||
|
* temperature
|
||||||
|
* temperature
|
||||||
|
)
|
||||||
|
return mx.sum(per_record * weights) / mx.sum(weights)
|
||||||
|
|
||||||
|
|
||||||
def _calibration(
|
def _calibration(
|
||||||
outputs: dict[str, np.ndarray],
|
outputs: dict[str, np.ndarray],
|
||||||
records: Sequence[dict[str, Any]],
|
records: Sequence[dict[str, Any]],
|
||||||
@@ -379,8 +412,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
encoded_validation = _encode_records(tokenizer, validation_records)
|
encoded_validation = _encode_records(tokenizer, validation_records)
|
||||||
train_targets = encode_targets(train_records)
|
train_targets = encode_targets(train_records)
|
||||||
validation_targets = encode_targets(validation_records)
|
validation_targets = encode_targets(validation_records)
|
||||||
|
validation_scorable = np.asarray(
|
||||||
|
[record["slice"] != "vague-eval" for record in validation_records],
|
||||||
|
dtype=np.bool_,
|
||||||
|
)
|
||||||
sample_weights = _sample_weights(train_records, args.hard_weight)
|
sample_weights = _sample_weights(train_records, args.hard_weight)
|
||||||
secondary_class_weights = _secondary_class_weights(train_records)
|
secondary_class_weights = _secondary_class_weights(train_records)
|
||||||
|
teacher_train_logits = None
|
||||||
|
teacher_validation_logits = None
|
||||||
|
if args.distillation_cache is not None:
|
||||||
|
teacher_train_logits, teacher_validation_logits = _teacher_cache(
|
||||||
|
args.distillation_cache.expanduser(),
|
||||||
|
train_records,
|
||||||
|
validation_records,
|
||||||
|
)
|
||||||
non_mixed = len(train_records) - int(train_targets.mixed.sum())
|
non_mixed = len(train_records) - int(train_targets.mixed.sum())
|
||||||
mixed_positive_weight = math.sqrt(
|
mixed_positive_weight = math.sqrt(
|
||||||
non_mixed / max(float(train_targets.mixed.sum()), 1.0)
|
non_mixed / max(float(train_targets.mixed.sum()), 1.0)
|
||||||
@@ -439,7 +484,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
mixed: Any,
|
mixed: Any,
|
||||||
difficulty: Any,
|
difficulty: Any,
|
||||||
weights: Any,
|
weights: Any,
|
||||||
) -> tuple[Any, Any, Any, Any, Any]:
|
teacher_logits: Any | None,
|
||||||
|
) -> tuple[Any, Any, Any, Any, Any, Any]:
|
||||||
output = model(input_ids=input_ids, attention_mask=attention_mask)
|
output = model(input_ids=input_ids, attention_mask=attention_mask)
|
||||||
primary_per_record = nn.losses.cross_entropy(
|
primary_per_record = nn.losses.cross_entropy(
|
||||||
output["purpose_logits"],
|
output["purpose_logits"],
|
||||||
@@ -447,7 +493,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
label_smoothing=args.label_smoothing,
|
label_smoothing=args.label_smoothing,
|
||||||
reduction="none",
|
reduction="none",
|
||||||
)
|
)
|
||||||
primary_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
|
primary_label_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
|
||||||
|
primary_distillation_loss = mx.zeros_like(primary_label_loss)
|
||||||
|
if teacher_logits is not None:
|
||||||
|
primary_distillation_loss = _distillation_loss(
|
||||||
|
mx,
|
||||||
|
output["purpose_logits"],
|
||||||
|
teacher_logits,
|
||||||
|
temperature=args.distillation_temperature,
|
||||||
|
weights=weights,
|
||||||
|
)
|
||||||
|
primary_loss = (
|
||||||
|
(1.0 - args.distillation_weight) * primary_label_loss
|
||||||
|
+ args.distillation_weight * primary_distillation_loss
|
||||||
|
)
|
||||||
|
|
||||||
secondary_per_record = nn.losses.cross_entropy(
|
secondary_per_record = nn.losses.cross_entropy(
|
||||||
output["secondary_logits"],
|
output["secondary_logits"],
|
||||||
@@ -482,23 +541,67 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
+ args.mixed_loss_weight * mixed_loss
|
+ args.mixed_loss_weight * mixed_loss
|
||||||
+ args.difficulty_loss_weight * difficulty_loss
|
+ args.difficulty_loss_weight * difficulty_loss
|
||||||
)
|
)
|
||||||
return total, primary_loss, secondary_loss, mixed_loss, difficulty_loss
|
return (
|
||||||
|
total,
|
||||||
|
primary_label_loss,
|
||||||
|
secondary_loss,
|
||||||
|
mixed_loss,
|
||||||
|
difficulty_loss,
|
||||||
|
primary_distillation_loss,
|
||||||
|
)
|
||||||
|
|
||||||
loss_and_grad = nn.value_and_grad(model, loss_function)
|
loss_and_grad = nn.value_and_grad(model, loss_function)
|
||||||
rng = np.random.default_rng(args.seed)
|
rng = np.random.default_rng(args.seed)
|
||||||
checkpoint_config = _checkpoint_config(source_config, variant)
|
checkpoint_config = _checkpoint_config(source_config, variant)
|
||||||
best_dir = output_dir / "model"
|
best_dir = output_dir / "model"
|
||||||
best_score = float("-inf")
|
|
||||||
best_metrics: dict[str, Any] | None = None
|
|
||||||
epochs_without_improvement = 0
|
epochs_without_improvement = 0
|
||||||
stopped_early = False
|
stopped_early = False
|
||||||
history: list[dict[str, Any]] = []
|
history: list[dict[str, Any]] = []
|
||||||
started = time.perf_counter()
|
started = time.perf_counter()
|
||||||
|
|
||||||
|
initial_outputs = _evaluate(
|
||||||
|
mx, model, encoded_validation, eval_batch_size
|
||||||
|
)
|
||||||
|
initial_metrics = multitask_metrics(initial_outputs, validation_records)
|
||||||
|
if teacher_validation_logits is not None:
|
||||||
|
initial_metrics["teacherAgreement"] = float(
|
||||||
|
np.mean(
|
||||||
|
initial_outputs["purpose_logits"][validation_scorable].argmax(
|
||||||
|
axis=-1
|
||||||
|
)
|
||||||
|
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
initial_metrics["epoch"] = 0
|
||||||
|
best_score = float(initial_metrics["selectionScore"])
|
||||||
|
best_metrics: dict[str, Any] = initial_metrics
|
||||||
|
_save_checkpoint(mx, model, tokenizer, best_dir, checkpoint_config)
|
||||||
|
write_json(
|
||||||
|
output_dir / "training-state.json",
|
||||||
|
{
|
||||||
|
"bestEpoch": 0,
|
||||||
|
"bestSelectionScore": best_score,
|
||||||
|
"elapsedSeconds": time.perf_counter() - started,
|
||||||
|
"complete": False,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"epoch 0: primary_accuracy={initial_metrics['primary']['accuracy']:.4%} "
|
||||||
|
f"hard_accuracy={initial_metrics['primaryHardSlice']['accuracy']:.4%} "
|
||||||
|
f"mixed_f1={initial_metrics['mixed']['f1']:.4%} "
|
||||||
|
f"selection_score={best_score:.4%}"
|
||||||
|
+ (
|
||||||
|
f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}"
|
||||||
|
if "teacherAgreement" in initial_metrics
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
for epoch in range(1, args.epochs + 1):
|
for epoch in range(1, args.epochs + 1):
|
||||||
epoch_started = time.perf_counter()
|
epoch_started = time.perf_counter()
|
||||||
model.train()
|
model.train()
|
||||||
running = np.zeros(5, dtype=np.float64)
|
running = np.zeros(6, dtype=np.float64)
|
||||||
permutation = rng.permutation(len(train_records))
|
permutation = rng.permutation(len(train_records))
|
||||||
for step, indexes in enumerate(
|
for step, indexes in enumerate(
|
||||||
_batch_indexes(
|
_batch_indexes(
|
||||||
@@ -509,6 +612,11 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
batch = _mlx_batch(
|
batch = _mlx_batch(
|
||||||
mx, encoded_train, train_targets, sample_weights, indexes
|
mx, encoded_train, train_targets, sample_weights, indexes
|
||||||
)
|
)
|
||||||
|
teacher_logits = (
|
||||||
|
mx.array(teacher_train_logits[indexes])
|
||||||
|
if teacher_train_logits is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
losses, gradients = loss_and_grad(
|
losses, gradients = loss_and_grad(
|
||||||
batch["input_ids"],
|
batch["input_ids"],
|
||||||
batch["attention_mask"],
|
batch["attention_mask"],
|
||||||
@@ -518,6 +626,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
batch["mixed"],
|
batch["mixed"],
|
||||||
batch["difficulty"],
|
batch["difficulty"],
|
||||||
batch["sample_weights"],
|
batch["sample_weights"],
|
||||||
|
teacher_logits,
|
||||||
)
|
)
|
||||||
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
|
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
|
||||||
optimizer.update(model, gradients)
|
optimizer.update(model, gradients)
|
||||||
@@ -531,7 +640,13 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
f"epoch {epoch} step {step}/{steps_per_epoch} "
|
f"epoch {epoch} step {step}/{steps_per_epoch} "
|
||||||
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
|
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
|
||||||
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
|
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
|
||||||
f"difficulty={mean[4]:.4f} "
|
f"difficulty={mean[4]:.4f}"
|
||||||
|
+ (
|
||||||
|
f" distillation={mean[5]:.4f}"
|
||||||
|
if teacher_train_logits is not None
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
+ " "
|
||||||
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
||||||
flush=True,
|
flush=True,
|
||||||
)
|
)
|
||||||
@@ -540,22 +655,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
mx, model, encoded_validation, eval_batch_size
|
mx, model, encoded_validation, eval_batch_size
|
||||||
)
|
)
|
||||||
metrics = multitask_metrics(outputs, validation_records)
|
metrics = multitask_metrics(outputs, validation_records)
|
||||||
|
if teacher_validation_logits is not None:
|
||||||
|
metrics["teacherAgreement"] = float(
|
||||||
|
np.mean(
|
||||||
|
outputs["purpose_logits"][validation_scorable].argmax(axis=-1)
|
||||||
|
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
|
||||||
|
)
|
||||||
|
)
|
||||||
metrics["epoch"] = epoch
|
metrics["epoch"] = epoch
|
||||||
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
|
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
|
||||||
history.append(metrics)
|
history.append(metrics)
|
||||||
score = float(metrics["selectionScore"])
|
score = float(metrics["selectionScore"])
|
||||||
secondary_macro = (
|
secondary_macro = (
|
||||||
metrics["secondary"]["macroRecall"]
|
metrics["secondary"]["supportedMacroRecall"]
|
||||||
if metrics["secondary"] is not None
|
if metrics["secondary"] is not None
|
||||||
else 0.0
|
else 0.0
|
||||||
)
|
)
|
||||||
print(
|
print(
|
||||||
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
|
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
|
||||||
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
|
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
|
||||||
f"secondary_macro_recall={secondary_macro:.4%} "
|
f"secondary_supported_macro_recall={secondary_macro:.4%} "
|
||||||
f"mixed_f1={metrics['mixed']['f1']:.4%} "
|
f"mixed_f1={metrics['mixed']['f1']:.4%} "
|
||||||
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
|
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
|
||||||
f"selection_score={score:.4%}",
|
f"selection_score={score:.4%}"
|
||||||
|
+ (
|
||||||
|
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
|
||||||
|
if "teacherAgreement" in metrics
|
||||||
|
else ""
|
||||||
|
),
|
||||||
flush=True,
|
flush=True,
|
||||||
)
|
)
|
||||||
improvement = score - best_score
|
improvement = score - best_score
|
||||||
@@ -588,9 +715,6 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
if best_metrics is None:
|
|
||||||
raise DataError("purpose-deep training did not produce a checkpoint")
|
|
||||||
|
|
||||||
# Release the optimizer graph before opening the selected checkpoint; base and
|
# Release the optimizer graph before opening the selected checkpoint; base and
|
||||||
# especially large should never hold two full optimizer states at calibration time.
|
# especially large should never hold two full optimizer states at calibration time.
|
||||||
del optimizer, loss_and_grad, model
|
del optimizer, loss_and_grad, model
|
||||||
@@ -632,6 +756,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"mixed": args.mixed_loss_weight,
|
"mixed": args.mixed_loss_weight,
|
||||||
"difficulty": args.difficulty_loss_weight,
|
"difficulty": args.difficulty_loss_weight,
|
||||||
},
|
},
|
||||||
|
"distillation": {
|
||||||
|
"cache": (
|
||||||
|
str(args.distillation_cache)
|
||||||
|
if args.distillation_cache is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
"weight": args.distillation_weight,
|
||||||
|
"temperature": args.distillation_temperature,
|
||||||
|
},
|
||||||
"secondaryClassWeights": {
|
"secondaryClassWeights": {
|
||||||
label: float(secondary_class_weights[index])
|
label: float(secondary_class_weights[index])
|
||||||
for index, label in enumerate(LABELS)
|
for index, label in enumerate(LABELS)
|
||||||
@@ -641,6 +774,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"batchSize": batch_size,
|
"batchSize": batch_size,
|
||||||
"learningRate": learning_rate,
|
"learningRate": learning_rate,
|
||||||
"bestValidationSelectionScore": best_score,
|
"bestValidationSelectionScore": best_score,
|
||||||
|
"initialValidation": initial_metrics,
|
||||||
"bestValidation": best_metrics,
|
"bestValidation": best_metrics,
|
||||||
"selectedValidation": calibrated,
|
"selectedValidation": calibrated,
|
||||||
"epochsCompleted": len(history),
|
"epochsCompleted": len(history),
|
||||||
@@ -707,6 +841,9 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
parser.add_argument("--secondary-loss-weight", type=float, default=0.25)
|
parser.add_argument("--secondary-loss-weight", type=float, default=0.25)
|
||||||
parser.add_argument("--mixed-loss-weight", type=float, default=0.25)
|
parser.add_argument("--mixed-loss-weight", type=float, default=0.25)
|
||||||
parser.add_argument("--difficulty-loss-weight", type=float, default=0.10)
|
parser.add_argument("--difficulty-loss-weight", type=float, default=0.10)
|
||||||
|
parser.add_argument("--distillation-cache", type=Path)
|
||||||
|
parser.add_argument("--distillation-weight", type=float, default=0.0)
|
||||||
|
parser.add_argument("--distillation-temperature", type=float, default=2.0)
|
||||||
parser.add_argument("--progress-steps", type=int, default=25)
|
parser.add_argument("--progress-steps", type=int, default=25)
|
||||||
parser.add_argument("--early-stopping-patience", type=int, default=1)
|
parser.add_argument("--early-stopping-patience", type=int, default=1)
|
||||||
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
|
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
|
||||||
@@ -750,6 +887,15 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||||||
):
|
):
|
||||||
if getattr(args, name) < 0:
|
if getattr(args, name) < 0:
|
||||||
parser.error(f"--{name.replace('_', '-')} must be non-negative")
|
parser.error(f"--{name.replace('_', '-')} must be non-negative")
|
||||||
|
if not 0 <= args.distillation_weight <= 1:
|
||||||
|
parser.error("--distillation-weight must be in [0, 1]")
|
||||||
|
if args.distillation_temperature <= 0:
|
||||||
|
parser.error("--distillation-temperature must be positive")
|
||||||
|
if (args.distillation_cache is None) != (args.distillation_weight == 0):
|
||||||
|
parser.error(
|
||||||
|
"--distillation-cache and a positive --distillation-weight "
|
||||||
|
"must be supplied together"
|
||||||
|
)
|
||||||
if not 0 < args.accepted_precision <= args.high_precision <= 1:
|
if not 0 < args.accepted_precision <= args.high_precision <= 1:
|
||||||
parser.error(
|
parser.error(
|
||||||
"confidence precision targets must satisfy 0 < accepted <= high <= 1"
|
"confidence precision targets must satisfy 0 < accepted <= high <= 1"
|
||||||
|
|||||||
Reference in New Issue
Block a user