Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 57 additions & 6 deletions src/bakery/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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")
Expand Down
47 changes: 47 additions & 0 deletions tests/test_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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() != ""
Loading