Skip to content
Open
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
5 changes: 3 additions & 2 deletions src/slime_bridge/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,9 @@ Slime calls one entry point, `generate_rollout_polar_async`, wired in via
(`group_id`, `policy_version`, `rollout_step`) onto every task, and keeps
async admission bounded to the current Slime rollout request;
- converts each Polar `Trajectory` back into Slime `Sample`s (one per trace,
grouped with Slime 0.3.0 `group_id` so all traces from a trajectory count
once), dropping empty or oversized traces;
grouped with the installed Slime trajectory-id field (`group_id` in the
v0.3.0 tag, `rollout_id` in ea9819f8/v0.3.1+) so all traces from a trajectory
count once), dropping empty or oversized traces;
- computes dynamic-trace leave-one-trajectory-out advantages and zeroes out
failed/aborted trajectories.

Expand Down
31 changes: 23 additions & 8 deletions src/slime_bridge/adapter.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
"""Convert Polar rollout results into Slime samples.

Every trace in ``Trajectory.traces`` becomes one Slime ``Sample``. All
samples produced from the same session share ``Sample.group_id`` so Slime
0.3.0's loss reducer counts the trajectory once even when it fans out into
multiple trace samples. Builders own trace curation and per-token loss masks
samples produced from the same session share Slime's trajectory id
(``group_id`` in the v0.3.0 tag, ``rollout_id`` in ea9819f8/v0.3.1+) so the
loss reducer counts the trajectory once even when it fans out into multiple
trace samples. Builders own trace curation and per-token loss masks
— the adapter does not infer trainable positions from bridge details. Traces
that lack training tokens are dropped and represented as fully masked samples
so callers can keep the rest of the group trainable.
Expand Down Expand Up @@ -39,7 +40,7 @@ def session_result_to_samples(
"""Convert one Polar session result into Slime samples — one per trace.

Every usable trace becomes an independent Sample sharing the same
``group_id`` key. Slime's loss reducer then averages all trace
trajectory-id key. Slime's loss reducer then averages all trace
contributions as one trajectory, while the reward post-processor can still
assign each trace its own advantage.

Expand Down Expand Up @@ -160,14 +161,15 @@ def _build_sample(
}
polar_metadata.update(_scheduler_metadata(result, trace))

return Sample(
return _make_sample(
Sample,
trajectory_id=index,
group_index=group_index,
index=index,
prompt=prompt_value,
tokens=prompt_ids + response_ids,
response=response_text,
response_length=len(response_ids),
group_id=index,
reward={reward_key: reward_value},
loss_mask=loss_mask,
rollout_log_probs=response_log_probs,
Expand Down Expand Up @@ -206,14 +208,15 @@ def _build_dummy_sample(
"placeholder": True,
}
polar_metadata.update(_scheduler_metadata(result, None))
return Sample(
return _make_sample(
Sample,
trajectory_id=index,
group_index=group_index,
index=index,
prompt="",
tokens=[0, 0],
response="",
response_length=1,
group_id=index,
reward={reward_key: 0.0},
loss_mask=[0],
rollout_log_probs=[0.0],
Expand All @@ -224,6 +227,18 @@ def _build_dummy_sample(
)


def _make_sample(Sample: Any, *, trajectory_id: int, **kwargs: Any) -> Any:
"""Use the trajectory-id field exposed by the installed Slime revision."""
fields = getattr(Sample, "__dataclass_fields__", {})
if "rollout_id" in fields: # ea9819f8 and v0.3.1+
kwargs["rollout_id"] = trajectory_id
elif "group_id" in fields: # v0.3.0 tag
kwargs["group_id"] = trajectory_id
else:
raise RuntimeError("Unsupported Slime Sample: missing rollout_id/group_id")
return Sample(**kwargs)


def _reward_value(trace: "Trace") -> float:
"""Read the reward the evaluator already placed on the trace.

Expand Down
9 changes: 6 additions & 3 deletions src/slime_bridge/reward_post_process.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,8 +8,9 @@

Adapter contract:
All Slime samples produced from the same Polar ``SessionResult`` share
``Sample.group_id``. Slime 0.3.0 uses that field to average all trace
contributions from one trajectory as one gradient unit.
one trajectory id: ``group_id`` in the v0.3.0 tag and ``rollout_id`` in
ea9819f8/v0.3.1+. Slime uses it to average all trace contributions from one
trajectory as one gradient unit.
"""

from __future__ import annotations
Expand Down Expand Up @@ -90,7 +91,9 @@ def post_process_rewards(

def _trajectory_key(sample: Any, sample_position: int) -> tuple[Any, tuple[Any, Any]]:
group_idx = _key_value(getattr(sample, "group_index", None), -1)
traj_idx = getattr(sample, "group_id", None)
traj_idx = getattr(sample, "rollout_id", None)
if traj_idx is None:
traj_idx = getattr(sample, "group_id", None)
if traj_idx is None:
traj_idx = getattr(sample, "index", None)
return group_idx, (group_idx, _key_value(traj_idx, sample_position))
Expand Down
26 changes: 20 additions & 6 deletions tests/slime_bridge/test_adapter.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

from dataclasses import dataclass
from enum import Enum

import pytest
Expand All @@ -11,6 +12,8 @@


class FakeSample:
__dataclass_fields__ = {"rollout_id": None}

class Status(str, Enum):
COMPLETED = "completed"
ABORTED = "aborted"
Expand All @@ -22,7 +25,7 @@ def __init__(
*,
group_index: int,
index: int,
group_id: int,
rollout_id: int,
prompt,
tokens: list[int],
response: str,
Expand All @@ -37,7 +40,7 @@ def __init__(
) -> None:
self.group_index = group_index
self.index = index
self.group_id = group_id
self.rollout_id = rollout_id
self.prompt = prompt
self.tokens = tokens
self.response = response
Expand All @@ -51,6 +54,17 @@ def __init__(
self.remove_sample = remove_sample


def test_make_sample_supports_slime_v030_group_id() -> None:
@dataclass
class V030Sample:
index: int
group_id: int

sample = adapter._make_sample(V030Sample, index=3, trajectory_id=7)

assert sample.group_id == 7


def _session_result(
*,
trace: Trace | None = None,
Expand Down Expand Up @@ -99,7 +113,7 @@ def test_session_result_to_samples_converts_trace_to_slime_like_sample(monkeypat
sample = samples[0]
assert sample.group_index == 11
assert sample.index == 2
assert sample.group_id == 2
assert sample.rollout_id == 2
assert sample.prompt == [{"role": "user", "content": "Say hi"}]
assert sample.tokens == [1, 2, 3, 4]
assert sample.response == "[assistant] Hi"
Expand All @@ -113,7 +127,7 @@ def test_session_result_to_samples_converts_trace_to_slime_like_sample(monkeypat
assert sample.metadata["polar"]["rollout_step"] == 7


def test_session_result_to_samples_shares_group_id_across_trace_siblings(monkeypatch) -> None:
def test_session_result_to_samples_shares_rollout_id_across_trace_siblings(monkeypatch) -> None:
monkeypatch.setattr(adapter, "_load_sample_type", lambda: FakeSample)
traces = [
Trace(
Expand All @@ -140,7 +154,7 @@ def test_session_result_to_samples_shares_group_id_across_trace_siblings(monkeyp

assert len(samples) == 2
assert [sample.index for sample in samples] == [7, 7]
assert [sample.group_id for sample in samples] == [7, 7]
assert [sample.rollout_id for sample in samples] == [7, 7]
assert [sample.metadata["polar"]["trace_index"] for sample in samples] == [0, 1]


Expand All @@ -160,7 +174,7 @@ def test_session_result_to_samples_emits_placeholder_when_trace_is_unusable(monk

assert len(samples) == 1
assert samples[0].remove_sample is True
assert samples[0].group_id == 2
assert samples[0].rollout_id == 2
assert samples[0].loss_mask == [0]
assert samples[0].metadata["polar"]["placeholder"] is True

Expand Down
38 changes: 22 additions & 16 deletions tests/slime_bridge/test_reward_post_process.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,23 +2,23 @@

from types import SimpleNamespace

from slime_bridge.reward_post_process import post_process_rewards
from slime_bridge.reward_post_process import _trajectory_key, post_process_rewards


class FakeSample:
def __init__(
self,
*,
group_index: int = 0,
group_id: int,
rollout_id: int,
reward: float,
status: str = "COMPLETED",
loss_mask: list[int] | None = None,
remove_sample: bool = False,
) -> None:
self.group_index = group_index
self.group_id = group_id
self.index = group_id
self.rollout_id = rollout_id
self.index = rollout_id
self.reward = {"score": reward}
self.status = status
self.loss_mask = [1] if loss_mask is None else loss_mask
Expand All @@ -40,11 +40,17 @@ def _args(**overrides):
return SimpleNamespace(**defaults)


def test_trajectory_key_supports_slime_v030_group_id() -> None:
sample = SimpleNamespace(group_index=2, group_id=7, index=99)

assert _trajectory_key(sample, 0) == (2, (2, 7))


def test_dynamic_trace_loo_keeps_per_trace_rewards() -> None:
samples = [
FakeSample(group_id=10, reward=2.0),
FakeSample(group_id=10, reward=4.0),
FakeSample(group_id=20, reward=10.0),
FakeSample(rollout_id=10, reward=2.0),
FakeSample(rollout_id=10, reward=4.0),
FakeSample(rollout_id=20, reward=10.0),
]

raw, rewards = post_process_rewards(_args(), samples)
Expand All @@ -55,9 +61,9 @@ def test_dynamic_trace_loo_keeps_per_trace_rewards() -> None:

def test_failed_trajectory_is_excluded_from_other_baselines() -> None:
samples = [
FakeSample(group_id=1, reward=2.0),
FakeSample(group_id=2, reward=10.0, status="FAILED"),
FakeSample(group_id=3, reward=6.0),
FakeSample(rollout_id=1, reward=2.0),
FakeSample(rollout_id=2, reward=10.0, status="FAILED"),
FakeSample(rollout_id=3, reward=6.0),
]

_, rewards = post_process_rewards(_args(), samples)
Expand All @@ -67,8 +73,8 @@ def test_failed_trajectory_is_excluded_from_other_baselines() -> None:

def test_single_valid_trajectory_uses_zero_baseline() -> None:
samples = [
FakeSample(group_id=1, reward=2.0),
FakeSample(group_id=1, reward=4.0),
FakeSample(rollout_id=1, reward=2.0),
FakeSample(rollout_id=1, reward=4.0),
]

_, rewards = post_process_rewards(_args(), samples)
Expand All @@ -78,8 +84,8 @@ def test_single_valid_trajectory_uses_zero_baseline() -> None:

def test_fully_masked_trajectory_does_not_enter_baseline() -> None:
samples = [
FakeSample(group_id=1, reward=2.0, loss_mask=[0], remove_sample=True),
FakeSample(group_id=2, reward=5.0),
FakeSample(rollout_id=1, reward=2.0, loss_mask=[0], remove_sample=True),
FakeSample(rollout_id=2, reward=5.0),
]

_, rewards = post_process_rewards(_args(), samples)
Expand All @@ -89,8 +95,8 @@ def test_fully_masked_trajectory_does_not_enter_baseline() -> None:

def test_disabled_normalization_returns_raw_rewards() -> None:
samples = [
FakeSample(group_id=1, reward=2.0),
FakeSample(group_id=2, reward=5.0),
FakeSample(rollout_id=1, reward=2.0),
FakeSample(rollout_id=2, reward=5.0),
]

raw, rewards = post_process_rewards(_args(rewards_normalization=False), samples)
Expand Down