your_name/jax_example

mean_confidence · authenticated

0.61

Input data provenance

unknown 100.0%

Evaluation source

@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)

Aggregation source

@measurement.aggregate
def aggregate(outputs):
    values = [value for batch in outputs for value in batch]
    return sum(values) / len(values)

Inputs and outputs

2 batches · page 1 of 1

Batch 1
Sample 1· 000000000000…

Input

Loading vector…

Output

Loading vector…
Sample 2· 000000000000…

Input

Loading vector…

Output

Loading vector…
Batch 2
Sample 1· 000000000000…

Input

Loading vector…

Output

Loading vector…
Sample 2· 000000000000…

Input

Loading vector…

Output

Loading vector…