From 02c05bbe9da7ca8185c83608fa024397f25f950d Mon Sep 17 00:00:00 2001 From: marksverdhei Date: Sun, 29 Mar 2026 00:47:46 +0000 Subject: [PATCH 1/3] perf: batch trajectory generation in a single model.generate call (issue #6 item 4) Previously training_step called _generate_trajectory once per (message, trajectory) pair, issuing N*T separate model.generate calls. This adds _generate_trajectories_batched which: 1. Repeats each user message num_trajectories times 2. Tokenizes all prompts with left-padding in one batch 3. Issues a single model.generate call for the whole batch 4. Decodes and pairs outputs back to their source messages Reduces GPU-CPU round trips from B*T to 1 per training step, which is meaningful when batch_size > 1 or num_trajectories > 1. Also adds 5 tests covering: pair types, message identity, num_trajectories repeat count, empty input, and empty-response filtering. Co-Authored-By: Claude Sonnet 4.6 --- src/bakery/trainer.py | 67 +++++++++++++++++++++++++++++++++++++++---- tests/test_trainer.py | 46 +++++++++++++++++++++++++++++ 2 files changed, 107 insertions(+), 6 deletions(-) diff --git a/src/bakery/trainer.py b/src/bakery/trainer.py index 29dbb43..b26edcb 100644 --- a/src/bakery/trainer.py +++ b/src/bakery/trainer.py @@ -145,6 +145,62 @@ 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) + + prompt_lengths = 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[prompt_lengths:], skip_special_tokens=True + ).strip() + if response: + results.append((msg, response)) + + return results + # -- Loss computation -- def compute_loss( @@ -414,12 +470,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..c590f56 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -283,3 +283,49 @@ 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() != "" From 6838ae92be8b68adb39ffe698934a99e77119fff Mon Sep 17 00:00:00 2001 From: marksverdhei Date: Sun, 29 Mar 2026 17:49:33 +0000 Subject: [PATCH 2/3] fix: rename prompt_lengths to padded_prompt_length for clarity The variable holds the padded sequence length (input_ids.shape[1]), not individual prompt lengths. Rename to avoid confusion with the _get_prompt_lengths method which returns per-message token counts. Co-Authored-By: Claude Opus 4.6 --- src/bakery/trainer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/bakery/trainer.py b/src/bakery/trainer.py index b26edcb..2b0fefb 100644 --- a/src/bakery/trainer.py +++ b/src/bakery/trainer.py @@ -177,7 +177,7 @@ def _generate_trajectories_batched( prompts, return_tensors="pt", padding=True ).to(self.model.device) - prompt_lengths = inputs["input_ids"].shape[1] + padded_prompt_length = inputs["input_ids"].shape[1] was_training = self.model.training self.model.eval() @@ -194,7 +194,7 @@ def _generate_trajectories_batched( results = [] for i, (msg, output_ids) in enumerate(zip(repeated_msgs, outputs)): response = self.processing_class.decode( - output_ids[prompt_lengths:], skip_special_tokens=True + output_ids[padded_prompt_length:], skip_special_tokens=True ).strip() if response: results.append((msg, response)) From 00d35bde5de78225a19a6fdf61d3b5287f1595e1 Mon Sep 17 00:00:00 2001 From: marksverdhei Date: Sun, 29 Mar 2026 17:50:45 +0000 Subject: [PATCH 3/3] style: apply ruff formatting --- src/bakery/trainer.py | 12 ++++-------- tests/test_trainer.py | 1 + 2 files changed, 5 insertions(+), 8 deletions(-) diff --git a/src/bakery/trainer.py b/src/bakery/trainer.py index 2b0fefb..d10dab1 100644 --- a/src/bakery/trainer.py +++ b/src/bakery/trainer.py @@ -166,16 +166,12 @@ def _generate_trajectories_batched( for msg in user_messages for _ in range(num_trajectories) ] - repeated_msgs = [ - 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) + inputs = self._tokenize(prompts, return_tensors="pt", padding=True).to( + self.model.device + ) padded_prompt_length = inputs["input_ids"].shape[1] diff --git a/tests/test_trainer.py b/tests/test_trainer.py index c590f56..e22298f 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -289,6 +289,7 @@ def test_prediction_step_sequential_eval_returns_triple(): # _generate_trajectories_batched # --------------------------------------------------------------------------- + def test_generate_trajectories_batched_returns_pairs(): """Batched generation returns (user_msg, response) pairs.""" trainer = _make_trainer(prompts=["Hello?"], responses=["Hi."])