From 9f8a4567b97049a3c4dd6a96c3a4530f7a19c148 Mon Sep 17 00:00:00 2001 From: Marcel Luethi Date: Tue, 29 Sep 2026 07:50:50 +0200 Subject: [PATCH] improve example in README --- README.md | 36 +++++++++++++++++++++--------------- mdocs/README.md | 36 +++++++++++++++++++++--------------- 2 files changed, 42 insertions(+), 30 deletions(-) diff --git a/README.md b/README.md index 2a88b5d7..04bff24a 100644 --- a/README.md +++ b/README.md @@ -32,29 +32,35 @@ JAX and einops, and efficient implementations of tensor operations using JAX as ```scala import dimwit.* +import dimwit.autodiff.Autodiff // Labels are simply Scala types trait Batch derives Label trait Feature derives Label +trait Hidden derives Label -// Create a 2D tensor with shape (3, 2), labeled with Batch and Feature -val t = Tensor( - Shape(Axis[Batch] -> 3, Axis[Feature] -> 2), -).fromArray( - Array( - 1.0f, 2.0f, - 3.0f, 4.0f, - 5.0f, 6.0f - ) +// Model parameters are plain case classes +case class Params(w: Tensor2[Feature, Hidden, Float32], b: Tensor1[Hidden, Float32]) + +// Write the model for a single example; axes are contracted by name, not position +def layer(p: Params)(x: Tensor1[Feature, Float32]): Tensor1[Hidden, Float32] = + (x.dot(Axis[Feature])(p.w) + p.b).tanh + +// ... and lift it to a whole batch with vmap +def loss(x: Tensor2[Batch, Feature, Float32])(p: Params): Tensor0[Float32] = + x.vmap(Axis[Batch])(layer(p)).pow(Tensor0(2.0f)).mean + +val x = Tensor(Shape(Axis[Batch] -> 32, Axis[Feature] -> 4)).fill(1.0f) +val params = Params( + Tensor(Shape(Axis[Feature] -> 4, Axis[Hidden] -> 8)).fill(0.1f), + Tensor(Shape(Axis[Hidden] -> 8)).fill(0.0f) ) -// Function to normalize a single feature vector -def normalize(x: Tensor1[Feature, Float32]) : Tensor1[Feature, Float32] = - (x -! x.mean) /! x.std +// Gradients have the same structure and types as the parameters +val grads: Params = Autodiff.grad(loss(x))(params).value -// Apply the normalization function across the Batch dimension -val normalized: Tensor2[Batch, Feature, Float32] = - t.vmap(Axis[Batch])(normalize) +// Mistakes are caught at compile time: +// x.dot(Axis[Hidden])(params.w) // error: Axis[Hidden] not found in Tensor[(Batch, Feature)] ``` See our [quickstart guide](docs/quickstart.md) for a more detailed introduction to the core concepts and API and diff --git a/mdocs/README.md b/mdocs/README.md index fffb63f2..7ee1d860 100644 --- a/mdocs/README.md +++ b/mdocs/README.md @@ -32,29 +32,35 @@ JAX and einops, and efficient implementations of tensor operations using JAX as ```scala mdoc:silent import dimwit.* +import dimwit.autodiff.Autodiff // Labels are simply Scala types trait Batch derives Label trait Feature derives Label +trait Hidden derives Label -// Create a 2D tensor with shape (3, 2), labeled with Batch and Feature -val t = Tensor( - Shape(Axis[Batch] -> 3, Axis[Feature] -> 2), -).fromArray( - Array( - 1.0f, 2.0f, - 3.0f, 4.0f, - 5.0f, 6.0f - ) +// Model parameters are plain case classes +case class Params(w: Tensor2[Feature, Hidden, Float32], b: Tensor1[Hidden, Float32]) + +// Write the model for a single example; axes are contracted by name, not position +def layer(p: Params)(x: Tensor1[Feature, Float32]): Tensor1[Hidden, Float32] = + (x.dot(Axis[Feature])(p.w) + p.b).tanh + +// ... and lift it to a whole batch with vmap +def loss(x: Tensor2[Batch, Feature, Float32])(p: Params): Tensor0[Float32] = + x.vmap(Axis[Batch])(layer(p)).pow(Tensor0(2.0f)).mean + +val x = Tensor(Shape(Axis[Batch] -> 32, Axis[Feature] -> 4)).fill(1.0f) +val params = Params( + Tensor(Shape(Axis[Feature] -> 4, Axis[Hidden] -> 8)).fill(0.1f), + Tensor(Shape(Axis[Hidden] -> 8)).fill(0.0f) ) -// Function to normalize a single feature vector -def normalize(x: Tensor1[Feature, Float32]) : Tensor1[Feature, Float32] = - (x -! x.mean) /! x.std +// Gradients have the same structure and types as the parameters +val grads: Params = Autodiff.grad(loss(x))(params).value -// Apply the normalization function across the Batch dimension -val normalized: Tensor2[Batch, Feature, Float32] = - t.vmap(Axis[Batch])(normalize) +// Mistakes are caught at compile time: +// x.dot(Axis[Hidden])(params.w) // error: Axis[Hidden] not found in Tensor[(Batch, Feature)] ``` See our [quickstart guide](docs/quickstart.md) for a more detailed introduction to the core concepts and API and