MCPcopy Create free account
hub / github.com/apple/axlearn / forward

Method forward

axlearn/common/evaler_test.py:131–136  ·  view source on GitHub ↗
(self, input_batch: NestedTensor)

Source from the content-addressed store, hash-verified

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:

Callers 2

test_alt_predictMethod · 0.45

Calls 2

WeightedSummaryClass · 0.90
add_summaryMethod · 0.45

Tested by

no test coverage detected