MCPcopy Create free account
hub / github.com/RL-Align/RL-Kernel / SyntheticRLKernelBatch

Class SyntheticRLKernelBatch

rl_engine/testing/rl_batch.py:13–84  ·  view source on GitHub ↗

Synthetic RL-shaped tensors shared by kernel tests and benchmarks.

Source from the content-addressed store, hash-verified

11
12@dataclass(frozen=True)
13class SyntheticRLKernelBatch:
14 """Synthetic RL-shaped tensors shared by kernel tests and benchmarks."""
15
16 input_ids: torch.Tensor
17 attention_mask: torch.Tensor
18 prompt_mask: torch.Tensor
19 completion_mask: torch.Tensor
20 token_ids: torch.Tensor
21 rewards: torch.Tensor
22 advantages: torch.Tensor
23 old_logps: torch.Tensor
24 ref_logps: torch.Tensor
25 valid_indices: torch.Tensor | None
26 metadata: dict[str, Any]
27
28 @property
29 def batch_size(self) -> int:
30 return int(self.input_ids.size(0))
31
32 @property
33 def total_seq_len(self) -> int:
34 return int(self.input_ids.size(1))
35
36 @property
37 def prompt_len(self) -> int:
38 return int(self.metadata["prompt_len"])
39
40 @property
41 def completion_len(self) -> int:
42 return int(self.metadata["completion_len"])
43
44 @property
45 def flat_completion_mask(self) -> torch.Tensor:
46 return self.completion_mask.reshape(-1)
47
48 @property
49 def flat_token_ids(self) -> torch.Tensor:
50 return self.token_ids.reshape(-1)
51
52 def dense_completion_token_ids(self) -> torch.Tensor:
53 return self.token_ids
54
55 def dense_completion_values(self, values: torch.Tensor) -> torch.Tensor:
56 expected_shape = (self.batch_size, self.completion_len)
57 if tuple(values.shape[:2]) != expected_shape:
58 raise ValueError(
59 f"expected leading shape {expected_shape}, got {tuple(values.shape[:2])}"
60 )
61 return values
62
63 def compact_completion_values(self, values: torch.Tensor) -> torch.Tensor:
64 dense = self.dense_completion_values(values)
65 return dense.reshape(-1, *dense.shape[2:])[self.flat_completion_mask]
66
67 def compact_token_ids(self) -> torch.Tensor:
68 return self.flat_token_ids[self.flat_completion_mask]
69
70 def benchmark_metadata(self) -> dict[str, Any]:

Callers 2

Calls

no outgoing calls

Tested by

no test coverage detected