diff --git a/src/bakery/trainer.py b/src/bakery/trainer.py index 29dbb43..d10dab1 100644 --- a/src/bakery/trainer.py +++ b/src/bakery/trainer.py @@ -145,6 +145,58 @@ def _generate_trajectory(self, user_message: str) -> str: ) return response.strip() + def _generate_trajectories_batched( + self, user_messages: list[str], num_trajectories: int + ) -> list[tuple[str, str]]: + """Generate multiple trajectories in a single batched model.generate call. + + Repeats each user message num_trajectories times, pads all prompts to + the same length using left-padding (so generated tokens align at the + right), runs a single model.generate, then decodes and pairs results. + + Returns: + List of (user_message, response) pairs, only for non-empty responses. + """ + if not user_messages: + return [] + + # Build one prompt per (message, trajectory) pair + prompts = [ + self._format_prompted(msg) + for msg in user_messages + for _ in range(num_trajectories) + ] + repeated_msgs = [msg for msg in user_messages for _ in range(num_trajectories)] + + with padding_side(self.processing_class, "left"): + inputs = self._tokenize(prompts, return_tensors="pt", padding=True).to( + self.model.device + ) + + padded_prompt_length = inputs["input_ids"].shape[1] + + was_training = self.model.training + self.model.eval() + + with torch.no_grad(): + with disable_adapters(self.model): + outputs = self.model.generate( + **inputs, generation_config=self.generation_config + ) + + if was_training: + self.model.train() + + results = [] + for i, (msg, output_ids) in enumerate(zip(repeated_msgs, outputs)): + response = self.processing_class.decode( + output_ids[padded_prompt_length:], skip_special_tokens=True + ).strip() + if response: + results.append((msg, response)) + + return results + # -- Loss computation -- def compute_loss( @@ -414,12 +466,11 @@ def training_step(self, model, inputs, num_items_in_batch=None) -> torch.Tensor: self.num_trajectories, ) - for user_msg in user_messages: - for _ in range(self.num_trajectories): - response = self._generate_trajectory(user_msg) - if response.strip(): - all_user_messages.append(user_msg) - all_responses.append(response) + for msg, resp in self._generate_trajectories_batched( + user_messages, self.num_trajectories + ): + all_user_messages.append(msg) + all_responses.append(resp) if not all_responses: logger.warning("No valid trajectories generated — returning zero loss") diff --git a/tests/test_trainer.py b/tests/test_trainer.py index f1dc148..e22298f 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -283,3 +283,50 @@ def test_prediction_step_sequential_eval_returns_triple(): assert result[1] is None and result[2] is None loss = result[0] assert loss.dim() == 0 + + +# --------------------------------------------------------------------------- +# _generate_trajectories_batched +# --------------------------------------------------------------------------- + + +def test_generate_trajectories_batched_returns_pairs(): + """Batched generation returns (user_msg, response) pairs.""" + trainer = _make_trainer(prompts=["Hello?"], responses=["Hi."]) + results = trainer._generate_trajectories_batched(["Hello?"], num_trajectories=1) + assert isinstance(results, list) + assert len(results) > 0 + for msg, resp in results: + assert isinstance(msg, str) + assert isinstance(resp, str) + + +def test_generate_trajectories_batched_msg_matches_input(): + """Each returned user_msg must be one of the original input messages.""" + trainer = _make_trainer(prompts=["Q1", "Q2"], responses=["A1", "A2"]) + results = trainer._generate_trajectories_batched(["Q1", "Q2"], num_trajectories=1) + returned_msgs = {msg for msg, _ in results} + assert returned_msgs.issubset({"Q1", "Q2"}) + + +def test_generate_trajectories_batched_num_trajectories(): + """With num_trajectories=2 each message appears up to 2 times.""" + trainer = _make_trainer(prompts=["Hi?"], responses=["Hello!"]) + results = trainer._generate_trajectories_batched(["Hi?"], num_trajectories=2) + # At most 2 results for 1 prompt * 2 trajectories + assert len(results) <= 2 + + +def test_generate_trajectories_batched_empty_messages(): + """Empty user_messages list returns empty results.""" + trainer = _make_trainer(prompts=["placeholder"], responses=["placeholder"]) + results = trainer._generate_trajectories_batched([], num_trajectories=2) + assert results == [] + + +def test_generate_trajectories_batched_filters_empty_responses(): + """Pairs with empty responses after strip() must be excluded.""" + trainer = _make_trainer(prompts=["placeholder"], responses=["placeholder"]) + results = trainer._generate_trajectories_batched(["Hi"], num_trajectories=1) + for _, resp in results: + assert resp.strip() != ""