Batch 1
Sample 1· 000000000000…
Input
Loading vector…
Output
Loading vector…
Sample 2· 000000000000…
Input
Loading vector…
Output
Loading vector…
mean_confidence · authenticated
@measurement.evaluate(params=0, batch=1)
def evaluate(params, batch):
logits = model.apply({"params": params}, batch)
return jax.nn.softmax(logits, axis=-1).max(axis=-1)
@measurement.aggregate
def aggregate(outputs):
values = [value for batch in outputs for value in batch]
return sum(values) / len(values)
2 batches · page 1 of 1