From 296d0008db94cd56e53281e8d1aa83c5f45c4123 Mon Sep 17 00:00:00 2001 From: marksverdhei Date: Thu, 19 Mar 2026 13:10:38 +0100 Subject: [PATCH] refactor: deduplicate trainer and clean up load_data API - Extract _prepare_inputs(), _per_sample_kl(), _validate_batch(), and _zero_loss() from compute_loss/prediction_step to eliminate ~60 lines of duplication - Change load_data() to accept explicit dataset/split kwargs, removing the type("_DC", ...) hack from cli.py Co-Authored-By: Claude Opus 4.6 (1M context) --- src/bakery/cli.py | 11 +- src/bakery/data.py | 28 +++- src/bakery/trainer.py | 299 ++++++++++++++++++++---------------------- 3 files changed, 169 insertions(+), 169 deletions(-) diff --git a/src/bakery/cli.py b/src/bakery/cli.py index 8ca0e4c..9b6f73f 100644 --- a/src/bakery/cli.py +++ b/src/bakery/cli.py @@ -245,15 +245,8 @@ def main(): if data_config.eval_dataset_split and data_config.dataset: print(f" Loading eval split: {data_config.eval_dataset_split}") eval_prompts, eval_responses = load_data( - type( - "_DC", - (), - { - "dataset": data_config.dataset, - "dataset_split": data_config.eval_dataset_split, - "training_prompts": None, - }, - )() + dataset=data_config.dataset, + dataset_split=data_config.eval_dataset_split, ) eval_dataset = create_dataset(eval_prompts, eval_responses) print(f" Eval samples: {len(eval_dataset)}") diff --git a/src/bakery/data.py b/src/bakery/data.py index a3218a4..6d5f9d2 100644 --- a/src/bakery/data.py +++ b/src/bakery/data.py @@ -91,19 +91,35 @@ def build_system_prompt( def load_data( - data_config: DataConfig, + data_config: Optional[DataConfig] = None, + *, + dataset: Optional[str] = None, + dataset_split: str = "train", + training_prompts: Optional[List[str]] = None, ) -> Tuple[List[str], Optional[List[str]]]: - """Load training data from config. + """Load training data from config or explicit parameters. + + Can be called with a DataConfig object (backward-compatible) or with + explicit keyword arguments for dataset/split/training_prompts. Returns: (prompts, responses) where responses is None if only prompts are available (triggering on-the-fly trajectory generation). """ - if data_config.dataset: - return load_dataset(data_config.dataset, data_config.dataset_split) + if data_config is not None: + dataset = dataset or data_config.dataset + dataset_split = ( + data_config.dataset_split + if dataset == data_config.dataset + else dataset_split + ) + training_prompts = training_prompts or data_config.training_prompts + + if dataset: + return load_dataset(dataset, dataset_split) - if data_config.training_prompts: - return data_config.training_prompts, None + if training_prompts: + return training_prompts, None raise ValueError( "No training data configured. Provide 'dataset' or 'training_prompts'." diff --git a/src/bakery/trainer.py b/src/bakery/trainer.py index 29dbb43..669af7a 100644 --- a/src/bakery/trainer.py +++ b/src/bakery/trainer.py @@ -145,32 +145,15 @@ def _generate_trajectory(self, user_message: str) -> str: ) return response.strip() - # -- Loss computation -- - - def compute_loss( - self, model, inputs, return_outputs=False, num_items_in_batch=None - ): - """Compute KL divergence loss with batched forward passes.""" - user_messages = inputs.get("user_messages", []) - responses = inputs.get("responses", []) + # -- Shared helpers -- - if not user_messages or not responses: - logger.warning( - "Batch has no user_messages or responses — returning zero loss" - ) - loss = torch.tensor(0.0, device=self.args.device, requires_grad=True) - return (loss, None) if return_outputs else loss - - pairs = [ - (msg, resp) for msg, resp in zip(user_messages, responses) if resp.strip() - ] - if not pairs: - logger.warning( - "All responses in batch are empty/whitespace — returning zero loss" - ) - loss = torch.tensor(0.0, device=self.args.device, requires_grad=True) - return (loss, None) if return_outputs else loss + def _prepare_inputs(self, pairs, model): + """Format, tokenize, and build forward-pass dicts for teacher/student. + Returns: + (teacher_inputs, student_inputs, teacher_fwd, student_fwd, + teacher_prompt_lengths, student_prompt_lengths) + """ teacher_texts, student_texts = [], [] teacher_prompt_lengths, student_prompt_lengths = [], [] @@ -204,75 +187,148 @@ def compute_loss( student_texts, return_tensors="pt", padding=True ).to(model.device) - teacher_fwd = dict( - input_ids=teacher_inputs["input_ids"], - attention_mask=teacher_inputs["attention_mask"], - ) - student_fwd = dict( - input_ids=student_inputs["input_ids"], - attention_mask=student_inputs["attention_mask"], - ) - # Some architectures (e.g. Gemma 3) require token_type_ids during training. - # The tokenizer may not return them, so create zeros if the model expects them. - if hasattr(model.config, "model_type") and model.config.model_type in ( - "gemma3", - ): - teacher_fwd["token_type_ids"] = torch.zeros_like( - teacher_inputs["input_ids"] - ) - student_fwd["token_type_ids"] = torch.zeros_like( - student_inputs["input_ids"] + def _make_fwd(tok_inputs): + fwd = dict( + input_ids=tok_inputs["input_ids"], + attention_mask=tok_inputs["attention_mask"], ) - elif "token_type_ids" in teacher_inputs: - teacher_fwd["token_type_ids"] = teacher_inputs["token_type_ids"] - student_fwd["token_type_ids"] = student_inputs["token_type_ids"] + # Some architectures (e.g. Gemma 3) require token_type_ids during training. + if hasattr(model.config, "model_type") and model.config.model_type in ( + "gemma3", + ): + fwd["token_type_ids"] = torch.zeros_like(tok_inputs["input_ids"]) + elif "token_type_ids" in tok_inputs: + fwd["token_type_ids"] = tok_inputs["token_type_ids"] + return fwd - with torch.no_grad(): - with disable_adapters(model): - teacher_outputs = model(**teacher_fwd) + teacher_fwd = _make_fwd(teacher_inputs) + student_fwd = _make_fwd(student_inputs) - student_outputs = model(**student_fwd) + return ( + teacher_inputs, + student_inputs, + teacher_fwd, + student_fwd, + teacher_prompt_lengths, + student_prompt_lengths, + ) - # With left-padding, each sequence has leading pad tokens that shift - # the real content rightward. Compute per-sequence padding offsets. + def _per_sample_kl( + self, + pairs, + teacher_logits, + student_logits, + teacher_inputs, + student_inputs, + teacher_prompt_lengths, + student_prompt_lengths, + ) -> list[torch.Tensor]: + """Compute per-sample KL divergence losses from aligned logit slices. + + Handles left-padding offsets and prompt-length masking. + """ t_seq_len = teacher_inputs["input_ids"].shape[1] s_seq_len = student_inputs["input_ids"].shape[1] t_real_lengths = teacher_inputs["attention_mask"].sum(dim=1) s_real_lengths = student_inputs["attention_mask"].sum(dim=1) losses = [] - for i in range(len(pairs)): t_pad = t_seq_len - t_real_lengths[i].item() s_pad = s_seq_len - s_real_lengths[i].item() t_start = int(t_pad) + teacher_prompt_lengths[i] s_start = int(s_pad) + student_prompt_lengths[i] - # Logits at position t predict token t+1, so shift back by 1 to get - # the logits that correspond to each response token. Slice to -1 - # because the last position predicts a token beyond the sequence. - t_logits = teacher_outputs.logits[i : i + 1, t_start - 1 : -1, :] - s_logits = student_outputs.logits[i : i + 1, s_start - 1 : -1, :] + # Logits at position t predict token t+1, so shift back by 1. + # Slice to -1 because the last position predicts beyond the sequence. + t_log = teacher_logits[i : i + 1, t_start - 1 : -1, :] + s_log = student_logits[i : i + 1, s_start - 1 : -1, :] t_mask = teacher_inputs["attention_mask"][i : i + 1, t_start:] s_mask = student_inputs["attention_mask"][i : i + 1, s_start:] - min_len = min(t_logits.shape[1], s_logits.shape[1]) + min_len = min(t_log.shape[1], s_log.shape[1]) if min_len == 0: continue - t_logits = t_logits[:, :min_len, :] - s_logits = s_logits[:, :min_len, :] + t_log = t_log[:, :min_len, :] + s_log = s_log[:, :min_len, :] mask = (t_mask[:, :min_len] * s_mask[:, :min_len]).float() - loss = compute_kl_divergence( - t_logits.detach(), s_logits, mask, self.kl_temperature + losses.append( + compute_kl_divergence(t_log.detach(), s_log, mask, self.kl_temperature) + ) + return losses + + def _zero_loss(self): + """Return a zero loss tensor that supports backward().""" + return torch.tensor(0.0, device=self.args.device, requires_grad=True) + + def _validate_batch(self, inputs): + """Extract and validate (user_message, response) pairs from a batch. + + Returns: + pairs or None if batch is empty/invalid (with warnings logged). + """ + user_messages = inputs.get("user_messages", []) + responses = inputs.get("responses", []) + + if not user_messages or not responses: + logger.warning( + "Batch has no user_messages or responses — returning zero loss" + ) + return None + + pairs = [ + (msg, resp) for msg, resp in zip(user_messages, responses) if resp.strip() + ] + if not pairs: + logger.warning( + "All responses in batch are empty/whitespace — returning zero loss" ) - losses.append(loss) + return None + + return pairs + + # -- Loss computation -- + + def compute_loss( + self, model, inputs, return_outputs=False, num_items_in_batch=None + ): + """Compute KL divergence loss with batched forward passes.""" + pairs = self._validate_batch(inputs) + if pairs is None: + loss = self._zero_loss() + return (loss, None) if return_outputs else loss + + ( + teacher_inputs, + student_inputs, + teacher_fwd, + student_fwd, + teacher_prompt_lengths, + student_prompt_lengths, + ) = self._prepare_inputs(pairs, model) + + with torch.no_grad(): + with disable_adapters(model): + teacher_outputs = model(**teacher_fwd) + + student_outputs = model(**student_fwd) + + losses = self._per_sample_kl( + pairs, + teacher_outputs.logits, + student_outputs.logits, + teacher_inputs, + student_inputs, + teacher_prompt_lengths, + student_prompt_lengths, + ) if not losses: logger.warning("No aligned logit pairs after slicing — returning zero loss") - zero = torch.tensor(0.0, device=self.args.device, requires_grad=True) + zero = self._zero_loss() return (zero, None) if return_outputs else zero total_loss = torch.stack(losses).mean() @@ -285,65 +341,24 @@ def prediction_step(self, model, inputs, prediction_loss_only, ignore_keys=None) student forward pass to halve peak VRAM usage. """ if not self.args.sequential_eval: - # Unwrap to bypass Accelerate's fp32 upcast wrapper raw = model.module if hasattr(model, "module") else model with torch.no_grad(): loss = self.compute_loss(raw, inputs) return (loss.detach(), None, None) # Sequential eval: teacher → offload logits to CPU → student - user_messages = inputs.get("user_messages", []) - responses = inputs.get("responses", []) - pairs = [(m, r) for m, r in zip(user_messages, responses) if r.strip()] - if not pairs: - return ( - torch.tensor(0.0, device=self.args.device, requires_grad=True), - None, - None, - ) - - teacher_texts, student_texts = [], [] - teacher_prompt_lengths, student_prompt_lengths = [], [] - for user_msg, response in pairs: - t_msgs = [ - {"role": "system", "content": self.system_prompt}, - {"role": "user", "content": user_msg}, - {"role": "assistant", "content": response}, - ] - s_msgs = [ - {"role": "user", "content": user_msg}, - {"role": "assistant", "content": response}, - ] - teacher_texts.append( - self.processing_class.apply_chat_template(t_msgs, tokenize=False) - ) - student_texts.append( - self.processing_class.apply_chat_template(s_msgs, tokenize=False) - ) - t_len, s_len = self._get_prompt_lengths(user_msg) - teacher_prompt_lengths.append(t_len) - student_prompt_lengths.append(s_len) - - with padding_side(self.processing_class, "left"): - teacher_inputs = self._tokenize( - teacher_texts, return_tensors="pt", padding=True - ).to(model.device) - student_inputs = self._tokenize( - student_texts, return_tensors="pt", padding=True - ).to(model.device) - - def _make_fwd(tok_inputs): - fwd = dict( - input_ids=tok_inputs["input_ids"], - attention_mask=tok_inputs["attention_mask"], - ) - if hasattr(model.config, "model_type") and model.config.model_type in ( - "gemma3", - ): - fwd["token_type_ids"] = torch.zeros_like(tok_inputs["input_ids"]) - elif "token_type_ids" in tok_inputs: - fwd["token_type_ids"] = tok_inputs["token_type_ids"] - return fwd + pairs = self._validate_batch(inputs) + if pairs is None: + return (self._zero_loss(), None, None) + + ( + teacher_inputs, + student_inputs, + teacher_fwd, + student_fwd, + teacher_prompt_lengths, + student_prompt_lengths, + ) = self._prepare_inputs(pairs, model) # Accelerate replaces model.forward with a wrapper that upcasts to fp32. # Bypass by calling the CLASS forward method directly. @@ -351,46 +366,22 @@ def _make_fwd(tok_inputs): fwd_fn = type(base).forward with torch.no_grad(): with disable_adapters(base): - teacher_logits = fwd_fn(base, **_make_fwd(teacher_inputs)).logits.cpu() + teacher_logits = fwd_fn(base, **teacher_fwd).logits.cpu() torch.cuda.empty_cache() - student_outputs = fwd_fn(base, **_make_fwd(student_inputs)) - - t_seq_len = teacher_inputs["input_ids"].shape[1] - s_seq_len = student_inputs["input_ids"].shape[1] - t_real_lengths = teacher_inputs["attention_mask"].sum(dim=1) - s_real_lengths = student_inputs["attention_mask"].sum(dim=1) - - losses = [] - for i in range(len(pairs)): - t_pad = t_seq_len - t_real_lengths[i].item() - s_pad = s_seq_len - s_real_lengths[i].item() - t_start = int(t_pad) + teacher_prompt_lengths[i] - s_start = int(s_pad) + student_prompt_lengths[i] - - t_logits = teacher_logits[i : i + 1, t_start - 1 : -1, :].to(model.device) - s_logits = student_outputs.logits[i : i + 1, s_start - 1 : -1, :] - t_mask = teacher_inputs["attention_mask"][i : i + 1, t_start:] - s_mask = student_inputs["attention_mask"][i : i + 1, s_start:] - - min_len = min(t_logits.shape[1], s_logits.shape[1]) - if min_len == 0: - continue - mask = (t_mask[:, :min_len] * s_mask[:, :min_len]).float() - losses.append( - compute_kl_divergence( - t_logits[:, :min_len, :], - s_logits[:, :min_len, :], - mask, - self.kl_temperature, - ) - ) + student_logits = fwd_fn(base, **student_fwd).logits + + losses = self._per_sample_kl( + pairs, + teacher_logits.to(model.device), + student_logits, + teacher_inputs, + student_inputs, + teacher_prompt_lengths, + student_prompt_lengths, + ) if not losses: - return ( - torch.tensor(0.0, device=self.args.device, requires_grad=True), - None, - None, - ) + return (self._zero_loss(), None, None) return (torch.stack(losses).mean().detach(), None, None) def training_step(self, model, inputs, num_items_in_batch=None) -> torch.Tensor: