Synthetic RL-shaped tensors shared by kernel tests and benchmarks.
| 11 | |
| 12 | @dataclass(frozen=True) |
| 13 | class 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]: |
no outgoing calls
no test coverage detected