Perform a forward pass and optionally a backward pass on a batch of data. Args: data: The input data for the forward pass, typically containing tensors and metadata. loss_function: The loss function to optimize. See `verl.workers.roles.utils.losses` for exam
(self, data: TensorDict, loss_function: Callable, forward_only=False)
| 96 | raise NotImplementedError |
| 97 | |
| 98 | def forward_backward_batch(self, data: TensorDict, loss_function: Callable, forward_only=False) -> Any: |
| 99 | """ |
| 100 | Perform a forward pass and optionally a backward pass on a batch of data. |
| 101 | |
| 102 | Args: |
| 103 | data: The input data for the forward pass, typically containing tensors and metadata. |
| 104 | loss_function: The loss function to optimize. See `verl.workers.roles.utils.losses` for examples. |
| 105 | forward_only: If True, perform only the forward pass. If False, perform forward and backward pass. |
| 106 | |
| 107 | Returns: |
| 108 | Any: The output of the forward pass, which can be used for loss computation or other purposes. |
| 109 | """ |
| 110 | raise NotImplementedError |
| 111 | |
| 112 | def train_batch(self, data: TensorDict, loss_function: Callable) -> Any: |
| 113 | """ |
no outgoing calls
no test coverage detected