Skip to content
Merged
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
14 changes: 7 additions & 7 deletions megatron/rl/agent/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,8 +46,8 @@ class GroupedRolloutRequest(Request):
class Rollout(AgentBaseModel):
"""Data for language-based Rollout."""

trajectory: str
prompt_length: int | None = None
trajectory: list[str]
prompt_length: list[int] | None = None
reward: float = None
env_id: str | None = None
problem_id: str | None = None
Expand All @@ -56,19 +56,19 @@ class Rollout(AgentBaseModel):
class TokenRollout(AgentBaseModel):
"""Tokenized representation of a language-based Rollout."""

trajectory: list[int]
trajectory: list[list[int]]
reward: list[float] | float
generation_mask: list[list[int]] | list[bool] | None = None
logprobs: list[float] | None = None
generation_mask: list[list[bool]] | None = None
logprobs: list[list[float]] | None = None
env_id: str | None = None
problem_id: str | None = None


class ContrastiveRollout(AgentBaseModel):
"""Contrastive/Preference data for language-based Rollout."""

chosen_trajectory: str
rejected_trajectory: str
chosen_trajectory: list[str]
rejected_trajectory: list[str]


class Head2HeadRolloutRequest(Request):
Expand Down
8 changes: 4 additions & 4 deletions megatron/rl/agent/reward_only_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,16 +104,16 @@ async def rollout_from_response(
for x in range(len(response.token_ids))
]
rollout = TokenRollout(
trajectory=response.token_ids,
trajectory=[response.token_ids],
reward=await self.get_reward(response_text, golden),
logprobs=logprobs,
generation_mask=generation_mask,
logprobs=[logprobs],
generation_mask=[generation_mask],
env_id=self.env_id,
problem_id=golden['problem_id'] if 'problem_id' in golden else None,
)
else:
rollout = Rollout(
trajectory=raw_text,
trajectory=[raw_text],
reward=await self.get_reward(response_text, golden),
env_id=self.env_id,
problem_id=golden['problem_id'] if 'problem_id' in golden else None,
Expand Down
Loading
Loading