(self, input_batch: NestedTensor)
| 129 | self.forward_dtypes = [] |
| 130 | |
| 131 | def forward(self, input_batch: NestedTensor) -> tuple[Tensor, NestedTensor]: |
| 132 | inputs = input_batch["inputs"] |
| 133 | self.forward_dtypes.append(inputs.dtype) |
| 134 | logits = self.linear(inputs) |
| 135 | self.add_summary("model_logits_sum", WeightedSummary(logits.sum(), logits.shape[0])) |
| 136 | return logits.mean(), {"logits": logits} |
| 137 | |
| 138 | # pylint: disable-next=no-self-use,unused-argument |
| 139 | def alt_predict(self, input_batch: NestedTensor, **kwargs) -> NestedTensor: |
no test coverage detected