From 7aea21e86ff31eaf74332c1e5169e884094e1ecd Mon Sep 17 00:00:00 2001 From: Kauna <16511995+klei22@users.noreply.github.com> Date: Tue, 9 Jun 2026 06:38:10 -0700 Subject: [PATCH] Add frozen lm_head KL training option --- model.py | 6 +- optimization_and_search/run_experiments.py | 6 + optimization_and_search/run_from_yaml.py | 36 ++- run_exploration_monitor.py | 4 + train.py | 266 +++++++++++++++++++-- train_args.py | 39 +++ 6 files changed, 328 insertions(+), 29 deletions(-) diff --git a/model.py b/model.py index 8643db0661..3041ac27de 100644 --- a/model.py +++ b/model.py @@ -379,7 +379,7 @@ def add_embedding_gaussian_noise(self, embeddings, iter_num=None): return embeddings + noise return embeddings - def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict=None, dataset_idx=None, loss_fn=None): + def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict=None, dataset_idx=None, loss_fn=None, return_hidden=False): if token_dict is not None: token_list = list(token_dict.values()) # If target_dict is None (typical for inference), set target_list = None @@ -531,6 +531,8 @@ def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict logits = [logit[:, [-1], :] for logit in logits] losses = None + if return_hidden: + raise ValueError("return_hidden is only supported for single-context forward passes.") return logits, losses else: @@ -633,6 +635,8 @@ def forward(self, idx, targets=None, iter_num=None, token_dict=None, target_dict loss = None + if return_hidden: + return logits, loss, x return logits, loss # ------------------------------------------------------------------ # LATENT-CHAINING diff --git a/optimization_and_search/run_experiments.py b/optimization_and_search/run_experiments.py index 1f98dae130..948ead58d6 100644 --- a/optimization_and_search/run_experiments.py +++ b/optimization_and_search/run_experiments.py @@ -40,6 +40,8 @@ "ln_f_cosine_95", "rankme", "areq", + "target_lm_head_kl_val", + "target_lm_head_kl_train", "zeus_total_energy_j", "zeus_total_time_s", "zeus_avg_power_w", @@ -582,6 +584,8 @@ def read_metrics(out_dir: str) -> dict: "ln_f_cosine_95", "rankme", "areq", + "target_lm_head_kl_val", + "target_lm_head_kl_train", ] casts = [ float, @@ -606,6 +610,8 @@ def read_metrics(out_dir: str) -> dict: float, float, float, + float, + float, ] if len(base_metric_keys) != len(casts): diff --git a/optimization_and_search/run_from_yaml.py b/optimization_and_search/run_from_yaml.py index aaaa114530..9c5a14cd23 100644 --- a/optimization_and_search/run_from_yaml.py +++ b/optimization_and_search/run_from_yaml.py @@ -24,9 +24,29 @@ METRICS_FILENAME = "best_val_loss_and_iter.txt" METRIC_KEYS = [ "best_val_loss", - "best_val_iter", + "best_val_iter", "best_tokens", "num_params", + "better_than_chance", + "btc_per_param", + "peak_torch_allocated_mb", + "peak_torch_reserved_mb", + "peak_process_gpu_mb", + "iter_latency_avg", + "zeus_best_train_step_energy_j", + "avg_top1_prob", + "avg_top1_correct", + "avg_target_rank", + "avg_target_left_prob", + "avg_target_prob", + "target_rank_95", + "left_prob_95", + "avg_ln_f_cosine", + "ln_f_cosine_95", + "rankme", + "areq", + "target_lm_head_kl_val", + "target_lm_head_kl_train", ] def _parse_override_args(arg_list: list[str] | None) -> dict: @@ -78,12 +98,14 @@ def read_metrics(out_dir: str) -> dict: line = path.read_text().strip() parts = [p.strip() for p in line.split(',')] - # Take only the first 4 values and cast them appropriately - if len(parts) < len(METRIC_KEYS): - raise ValueError(f"Expected at least {len(METRIC_KEYS)} metrics, got {len(parts)}") - - casts = [float, int, int, int] - return {k: typ(v) for k, typ, v in zip(METRIC_KEYS, casts, parts[:len(METRIC_KEYS)])} + if len(parts) < 4: + raise ValueError(f"Expected at least 4 metrics, got {len(parts)}") + + casts = [float, int, int, int] + [float] * (len(METRIC_KEYS) - 4) + metrics = {} + for key, typ, value in zip(METRIC_KEYS, casts, parts[:len(METRIC_KEYS)]): + metrics[key] = float("nan") if value == "" else typ(value) + return metrics def completed_runs(log_file: Path) -> set[str]: diff --git a/run_exploration_monitor.py b/run_exploration_monitor.py index ec570233f7..1dde6ed4e7 100644 --- a/run_exploration_monitor.py +++ b/run_exploration_monitor.py @@ -203,6 +203,8 @@ def on_mount(self) -> None: "ln_f_cosine_95", "rankme", "areq", + "target_lm_head_kl_val", + "target_lm_head_kl_train", "zeus_total_energy_j", "zeus_total_time_s", "zeus_avg_power_w", @@ -321,6 +323,8 @@ def get_cell(self, entry: Dict, col_name: str): "ln_f_cosine_95", "rankme", "areq", + "target_lm_head_kl_val", + "target_lm_head_kl_train", "zeus_total_energy_j", "zeus_total_time_s", "zeus_avg_power_w", diff --git a/train.py b/train.py index d326db89c4..0ff07611b3 100644 --- a/train.py +++ b/train.py @@ -76,6 +76,7 @@ # Torch import torch import torch.onnx +import torch.nn as nn import torch.nn.functional as F from torch.distributed import destroy_process_group, init_process_group from torch.nn.parallel import DistributedDataParallel as DDP @@ -207,6 +208,11 @@ def __init__(self, args, model_group, training_group, logging_group): self.distillation_weight = getattr(self.args, "distillation_weight", 1.0) self.teacher_model = None self.latest_distillation_loss = float('nan') + self.target_lm_head = None + self.target_lm_head_softcap = None + self.latest_target_lm_head_kl = float('nan') + self.latest_target_lm_head_kl_val = float('nan') + self.latest_target_lm_head_kl_train_eval = float('nan') if self.distillation_loss_fn is not None and self.args.training_mode == 'multicontext': raise ValueError("Knowledge distillation is not supported with multicontext training mode.") @@ -431,6 +437,7 @@ def setup(self): self.scheduler = self.create_scheduler() self._initialize_teacher_if_needed() + self._initialize_target_lm_head_if_needed() if self.args.block_size < self.model.config.block_size: self.model.crop_block_size(self.args.block_size) @@ -543,6 +550,113 @@ def _initialize_teacher_if_needed(self): print(f"Loaded teacher checkpoint from {expanded}") + def _resolve_checkpoint_path(self, path_value): + expanded = os.path.expanduser(path_value) + if not os.path.exists(expanded): + candidate = os.path.join(self.args.out_dir, path_value) + if os.path.exists(candidate): + expanded = candidate + return expanded + + + def _target_lm_head_key(self, checkpoint_model_args, dataset_idx=None): + if checkpoint_model_args.get('multidataset_wte') or checkpoint_model_args.get('multicontext'): + idx = dataset_idx + if idx is None: + idx = getattr(self.args, 'target_lm_head_dataset_idx', None) + if idx is None: + idx = 0 + return f"transformer.lm_head_{idx}.weight", idx + return 'lm_head.weight', None + + + def _initialize_target_lm_head_if_needed(self): + mode = getattr(self.args, 'target_lm_head_kl_mode', 'off') + ckpt_path = getattr(self.args, 'target_lm_head_ckpt', None) + if mode == 'off': + if ckpt_path: + raise ValueError( + "--target_lm_head_ckpt requires --target_lm_head_kl_mode to be 'loss' or 'monitor'." + ) + return + if self.args.training_mode == 'multicontext': + raise ValueError("Frozen lm_head KL is not supported with multicontext training mode.") + if not ckpt_path: + raise ValueError("--target_lm_head_kl_mode requires --target_lm_head_ckpt.") + if self.args.target_lm_head_kl_temperature <= 0: + raise ValueError("target_lm_head_kl_temperature must be positive.") + if self.args.target_lm_head_kl_eps <= 0: + raise ValueError("target_lm_head_kl_eps must be positive.") + + expanded = self._resolve_checkpoint_path(ckpt_path) + checkpoint = torch.load(expanded, map_location='cpu') + checkpoint_model_args = checkpoint.get('model_args') + if checkpoint_model_args is None: + raise ValueError("Target lm_head checkpoint does not contain 'model_args'.") + + requested_idx = getattr(self.args, 'target_lm_head_dataset_idx', None) + key, resolved_idx = self._target_lm_head_key(checkpoint_model_args, requested_idx) + state_dict = checkpoint['model'] + clean_state = {} + for state_key, value in state_dict.items(): + if state_key.startswith('_orig_mod.'): + state_key = state_key[len('_orig_mod.'):] + clean_state[state_key] = value + + if key not in clean_state: + fallback_keys = ['lm_head.weight', 'transformer.wte.weight'] + if resolved_idx is not None: + fallback_keys.insert(0, f'transformer.wte_{resolved_idx}.weight') + for fallback in fallback_keys: + if fallback in clean_state: + key = fallback + break + else: + available = sorted( + k for k in clean_state + if k == 'lm_head.weight' or k == 'transformer.wte.weight' or '.lm_head_' in k or '.wte_' in k + ) + raise ValueError( + f"Could not find {key!r} in target lm_head checkpoint. Available lm_head/wte keys: {available}" + ) + + weight = clean_state[key].detach().float() + target_head = nn.Linear(weight.shape[1], weight.shape[0], bias=False) + target_head.weight.data.copy_(weight) + target_head.to(self.device) + target_head.eval() + for param in target_head.parameters(): + param.requires_grad_(False) + + self.target_lm_head_softcap = checkpoint_model_args.get('final_logit_softcapping') + self.target_lm_head = target_head + if self.master_process: + idx_msg = f" (dataset index {resolved_idx})" if resolved_idx is not None else "" + print(f"Loaded frozen target lm_head{idx_msg} from {expanded} using key {key}") + + + def _target_lm_head_kl_loss(self, student_logits, hidden): + if self.target_lm_head is None: + return None + if hidden.size(-1) != self.target_lm_head.in_features: + raise ValueError( + f"Frozen lm_head expects hidden size {self.target_lm_head.in_features}, got {hidden.size(-1)}." + ) + temperature = self.args.target_lm_head_kl_temperature + eps = self.args.target_lm_head_kl_eps + with torch.no_grad(): + target_logits = self.target_lm_head(hidden).to(student_logits.dtype) + if self.target_lm_head_softcap is not None: + target_logits = torch.tanh(target_logits / self.target_lm_head_softcap) * self.target_lm_head_softcap + student_log_probs = F.log_softmax(student_logits.float() / temperature, dim=-1) + target_probs = F.softmax(target_logits.float() / temperature, dim=-1) + target_probs = target_probs.clamp_min(eps) + target_probs = target_probs / target_probs.sum(dim=-1, keepdim=True) + kl = F.kl_div(student_log_probs, target_probs, reduction='none').sum(dim=-1) + kl = torch.clamp(kl, min=0.0) + return (kl.mean() * (temperature ** 2) + (student_logits.sum() * 0.0)).to(student_logits.dtype) + + def create_optimizer(self): optimizer_key = self.args.optimizer @@ -1006,6 +1120,7 @@ def estimate_loss(self): for dataset in self.args.dataset_list: print(f"Calculating loss for dataset: {dataset}") dataset_losses = {'train': torch.zeros(self.args.eval_iters), 'val': torch.zeros(self.args.eval_iters)} + dataset_lm_head_kl_losses = {'train': torch.zeros(self.args.eval_iters), 'val': torch.zeros(self.args.eval_iters)} top1_probs, top1_corrects, target_ranks, target_probs, target_left_probs, left_inclusive_probs = [], [], [], [], [], [] ln_f_cosines = [] rankme_vectors = [] if compute_rankme else None @@ -1018,15 +1133,31 @@ def estimate_loss(self): ) with self.ctx: idx = self.args.dataset_list.index(dataset) - logits, loss = self.model( - X, - Y, - iter_num=self.iter_num, - dataset_idx=idx if self.args.multidataset_wte else None, - loss_fn=self.loss_fn, - ) + if self.target_lm_head is not None: + logits, loss, hidden = self.model( + X, + Y, + iter_num=self.iter_num, + dataset_idx=idx if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + return_hidden=True, + ) + else: + logits, loss = self.model( + X, + Y, + iter_num=self.iter_num, + dataset_idx=idx if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + ) + hidden = None handle.remove() dataset_losses[split][k] = loss.item() + if self.target_lm_head is not None and hidden is not None: + kl_loss = self._target_lm_head_kl_loss(logits, hidden) + dataset_lm_head_kl_losses[split][k] = kl_loss.item() + else: + dataset_lm_head_kl_losses[split][k] = float('nan') if split == 'val': probs = F.softmax(logits, dim=-1) top1_prob, top1_idx = probs.max(dim=-1) @@ -1064,6 +1195,10 @@ def estimate_loss(self): 'train_std': dataset_losses['train'].std(), 'val': dataset_losses['val'].mean(), 'val_std': dataset_losses['val'].std(), + 'target_lm_head_kl_train': self._nanmean(dataset_lm_head_kl_losses['train']), + 'target_lm_head_kl_val': self._nanmean(dataset_lm_head_kl_losses['val']), + 'target_lm_head_kl_train_std': self._nanstd(dataset_lm_head_kl_losses['train']), + 'target_lm_head_kl_val_std': self._nanstd(dataset_lm_head_kl_losses['val']), 'top1_prob': torch.cat(top1_probs).mean() if top1_probs else torch.tensor(float('nan')), 'top1_correct': torch.cat(top1_corrects).mean() if top1_corrects else torch.tensor(float('nan')), 'target_rank': torch.cat(target_ranks).mean() if target_ranks else torch.tensor(float('nan')), @@ -1090,6 +1225,10 @@ def estimate_loss(self): out['left_prob_95'] = out['datasets'][self.args.dataset]['left_prob_95'] out['ln_f_cosine'] = out['datasets'][self.args.dataset]['ln_f_cosine'] out['ln_f_cosine_95'] = out['datasets'][self.args.dataset]['ln_f_cosine_95'] + out['target_lm_head_kl_train'] = out['datasets'][self.args.dataset]['target_lm_head_kl_train'] + out['target_lm_head_kl_val'] = out['datasets'][self.args.dataset]['target_lm_head_kl_val'] + out['target_lm_head_kl_train_std'] = out['datasets'][self.args.dataset]['target_lm_head_kl_train_std'] + out['target_lm_head_kl_val_std'] = out['datasets'][self.args.dataset]['target_lm_head_kl_val_std'] if compute_rankme: out['rankme'] = out['datasets'][self.args.dataset]['rankme'] out['areq'] = out['datasets'][self.args.dataset]['areq'] @@ -1143,6 +1282,7 @@ def estimate_loss(self): # Default behavior for a single dataset for split in ['train', 'val']: losses = torch.zeros(self.args.eval_iters) + lm_head_kl_losses = torch.zeros(self.args.eval_iters) top1_probs, top1_corrects, target_ranks, target_probs, target_left_probs, left_inclusive_probs = [], [], [], [], [], [] ln_f_cosines = [] rankme_vectors = [] if compute_rankme else None @@ -1153,15 +1293,31 @@ def estimate_loss(self): lambda _m, _i, o: ln_f_out.append(o.detach()) ) with self.ctx: - logits, loss = self.model( - X, - Y, - iter_num=self.iter_num, - dataset_idx=0 if self.args.multidataset_wte else None, - loss_fn=self.loss_fn, - ) + if self.target_lm_head is not None: + logits, loss, hidden = self.model( + X, + Y, + iter_num=self.iter_num, + dataset_idx=0 if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + return_hidden=True, + ) + else: + logits, loss = self.model( + X, + Y, + iter_num=self.iter_num, + dataset_idx=0 if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + ) + hidden = None handle.remove() losses[k] = loss.item() + if self.target_lm_head is not None and hidden is not None: + kl_loss = self._target_lm_head_kl_loss(logits, hidden) + lm_head_kl_losses[k] = kl_loss.item() + else: + lm_head_kl_losses[k] = float('nan') if split == 'val': probs = F.softmax(logits, dim=-1) top1_prob, top1_idx = probs.max(dim=-1) @@ -1191,6 +1347,8 @@ def estimate_loss(self): ) out[split] = losses.mean() out[split + "_std"] = losses.std() + out[f'target_lm_head_kl_{split}'] = self._nanmean(lm_head_kl_losses) + out[f'target_lm_head_kl_{split}_std'] = self._nanstd(lm_head_kl_losses) if split == 'val': out['top1_prob'] = torch.cat(top1_probs).mean() if top1_probs else torch.tensor(float('nan')) out['top1_correct'] = torch.cat(top1_corrects).mean() if top1_corrects else torch.tensor(float('nan')) @@ -1376,6 +1534,20 @@ def _log_bit_metrics(self, target_dataset: str, tokens_trained: float) -> None: f"{target_dataset}/bit_loss_penalty_tokens", penalty_term, tokens_trained ) + def _nanmean(self, values: torch.Tensor) -> torch.Tensor: + finite = torch.isfinite(values) + if not finite.any(): + return values.new_tensor(float('nan')) + return values[finite].mean() + + + def _nanstd(self, values: torch.Tensor) -> torch.Tensor: + finite = torch.isfinite(values) + if not finite.any(): + return values.new_tensor(float('nan')) + return values[finite].std() + + def _safe_better_than_chance(self, vocab_size: float, loss_value: float) -> float: """Return vocab_size / exp(loss) without raising overflow for huge losses.""" if not math.isfinite(loss_value): @@ -1456,6 +1628,18 @@ def log_metrics(self, losses, running_mfu, epoch, tokens_trained, target_dataset self.iter_num, ) + if 'target_lm_head_kl_val' in losses: + self.writer.add_scalar( + f"{target_dataset}/target_lm_head_kl_val", + losses['target_lm_head_kl_val'], + self.iter_num, + ) + self.writer.add_scalar( + f"{target_dataset}/target_lm_head_kl_train_eval", + losses['target_lm_head_kl_train'], + self.iter_num, + ) + if 'top1_prob' in losses: self.writer.add_scalar(f"{target_dataset}/avg_top1_prob", losses['top1_prob'], self.iter_num) self.writer.add_scalar(f"{target_dataset}/avg_top1_correct", losses['top1_correct'], self.iter_num) @@ -1572,6 +1756,13 @@ def log_metrics_non_validation(self, loss_training, running_mfu, epoch, tokens_t self.iter_num, ) + if not math.isnan(self.latest_target_lm_head_kl): + self.writer.add_scalar( + f"{target_dataset}/target_lm_head_kl_train_step", + self.latest_target_lm_head_kl, + self.iter_num, + ) + if self.args.log_grad_norm: self.writer.add_scalar(f"{target_dataset}/grad_norm_iters", self.grad_norm, self.iter_num) self.writer.add_scalar(f"{target_dataset}/grad_norm_tokens", self.grad_norm, tokens_trained) @@ -1795,6 +1986,8 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s self.latest_left_prob_95 = losses.get('left_prob_95', float('nan')) self.latest_ln_f_cosine = losses.get('ln_f_cosine', float('nan')) self.latest_ln_f_cosine_95 = losses.get('ln_f_cosine_95', float('nan')) + self.latest_target_lm_head_kl_val = self._to_scalar(losses.get('target_lm_head_kl_val', float('nan'))) + self.latest_target_lm_head_kl_train_eval = self._to_scalar(losses.get('target_lm_head_kl_train', float('nan'))) self.latest_rankme = self._to_scalar(losses.get('rankme', float('nan'))) self.latest_areq = self._to_scalar(losses.get('areq', float('nan'))) @@ -1831,6 +2024,8 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s log_message+=f", btc_val_per_param {(better_than_chance/self.model.num_param):.2e}" log_message+=f", val loss {dataset_losses['val']:.4f}" log_message+=f", val_stdev {dataset_losses['val_std']:.4f}" + if 'target_lm_head_kl_val' in dataset_losses: + log_message+=f", lm_head_kl_val {dataset_losses['target_lm_head_kl_val']:.4f}" if self.args.gns_type is not None: log_message+=f", gns {self.gns:.2f}" log_message+=f", lr {self.lr:.4f}" @@ -1862,6 +2057,8 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s log_message+=f", btc_val_per_param {(better_than_chance/self.model.num_param):.2e}" log_message+=f", val loss {losses['val']:.4f}" log_message+=f", val_stdev {losses['val_std']:.4f}" + if 'target_lm_head_kl_val' in losses: + log_message+=f", lm_head_kl_val {losses['target_lm_head_kl_val']:.4f}" if self.args.gns_type is not None: log_message+=f", gns {self.gns:.2f}" log_message+=f", batch_size {self.args.batch_size}" @@ -1916,6 +2113,8 @@ def run_validation_step(self, running_mfu, current_epoch, current_dataset, num_s f"{self.latest_ln_f_cosine_95:.6f}", f"{self.latest_rankme:.6f}", f"{self.latest_areq:.6f}", + f"{self.latest_target_lm_head_kl_val:.6f}", + f"{self.latest_target_lm_head_kl_train_eval:.6f}", f"{self.latest_overall_weight_stats['stdev']:.6f}", f"{self.latest_overall_weight_stats['kurtosis']:.6f}", f"{self.latest_overall_weight_stats['max']:.6f}", @@ -1973,6 +2172,8 @@ def log_training_step(self, lossf, training_losses, running_mfu, current_epoch, else: better_than_chance = self._safe_better_than_chance(self.model_args['vocab_size'], lossf) log_message+= f", loss {lossf:.4f}" + if not math.isnan(self.latest_target_lm_head_kl): + log_message+= f", lm_head_kl {self.latest_target_lm_head_kl:.4f}" if self.args.log_btc_train: log_message+=f", btc_train {better_than_chance:.2e}" if self.args.log_btc_per_param: @@ -2133,13 +2334,24 @@ def train(self): loss = sum(training_losses) / len(training_losses) else: idx_ds = self.args.dataset_list.index(current_dataset) if self.args.dataset_list else None - logits, loss = self.model( - self.X, - targets=self.Y, - iter_num=self.iter_num, - dataset_idx=idx_ds if self.args.multidataset_wte else None, - loss_fn=self.loss_fn, - ) + if self.target_lm_head is not None: + logits, loss, hidden = self.model( + self.X, + targets=self.Y, + iter_num=self.iter_num, + dataset_idx=idx_ds if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + return_hidden=True, + ) + else: + logits, loss = self.model( + self.X, + targets=self.Y, + iter_num=self.iter_num, + dataset_idx=idx_ds if self.args.multidataset_wte else None, + loss_fn=self.loss_fn, + ) + hidden = None if hasattr(self.optimizer, "set_entropy") and not isinstance(logits, (list, tuple)): with torch.no_grad(): @@ -2148,6 +2360,18 @@ def train(self): ent = ent / math.log(logits.size(-1)) self.optimizer.set_entropy(float(ent)) + target_lm_head_component = None + if self.target_lm_head is not None and hidden is not None: + target_lm_head_component = self._target_lm_head_kl_loss(logits, hidden) + if self.args.target_lm_head_kl_mode == 'loss': + loss = self.args.target_lm_head_kl_weight * target_lm_head_component + + self.latest_target_lm_head_kl = ( + float(target_lm_head_component.detach().float().item()) + if target_lm_head_component is not None + else float('nan') + ) + distill_component = None if ( self.teacher_model is not None diff --git a/train_args.py b/train_args.py index 7f7b2ace67..5520cabd2f 100644 --- a/train_args.py +++ b/train_args.py @@ -226,6 +226,45 @@ def parse_args(): help='Numerical stability epsilon for distillation losses.', ) + # Frozen LM-head KL options + training_group.add_argument( + '--target_lm_head_ckpt', + type=str, + default=None, + help='Path to a checkpoint whose frozen lm_head is used as the KL target.', + ) + training_group.add_argument( + '--target_lm_head_kl_mode', + type=str, + default='off', + choices=['off', 'loss', 'monitor'], + help="Use the frozen lm_head KL as the training loss ('loss'), compute it only for logging ('monitor'), or disable it.", + ) + training_group.add_argument( + '--target_lm_head_kl_weight', + type=float, + default=1.0, + help='Scaling factor applied when target_lm_head_kl_mode=loss.', + ) + training_group.add_argument( + '--target_lm_head_kl_temperature', + type=float, + default=1.0, + help='Temperature used for the frozen lm_head KL calculation.', + ) + training_group.add_argument( + '--target_lm_head_kl_eps', + type=float, + default=1e-8, + help='Numerical stability epsilon for the frozen lm_head KL calculation.', + ) + training_group.add_argument( + '--target_lm_head_dataset_idx', + type=int, + default=None, + help='Optional lm_head_N index to load from a multidataset/multicontext checkpoint. Defaults to the current dataset index when available, else 0.', + ) + # Sample args training_group.add_argument('--max_sample_tokens', default=None, type=int, help="If set, maximum number of tokens to sample and print after each validation loss") training_group.add_argument('--sample_each_eval', default=False, action=argparse.BooleanOptionalAction, help="Produce sample even if the validation loss did not improve. Allows for testing what overtraining looks like.")