Guides

Train a model

The Module API keeps the training loop explicit: clear gradients, run forward, backpropagate, update parameters, and release the batch graph.

Model definition

The MNIST example uses a 784 → 128 → 10 MLP with ReLU and Dropout. This snippet follows the source model; it omits only the seeded weight options.

MNIST modelexamples/mnist/model.ts
export class SimpleMLP extends nn.Module {
  private readonly flatten: nn.Flatten;
  private readonly fc1: nn.Linear;
  private readonly relu: nn.ReLU;
  private readonly dropout: nn.Dropout;
  private readonly fc2: nn.Linear;

  constructor() {
    super();
    this.flatten = this.register("flatten", new nn.Flatten());
    this.fc1 = this.register("fc1", new nn.Linear(784, 128));
    this.relu = this.register("relu", new nn.ReLU());
    this.dropout = this.register("dropout", new nn.Dropout(0.2));
    this.fc2 = this.register("fc2", new nn.Linear(128, 10));
  }

  forward(input: nn.Tensor): nn.Tensor {
    return this.fc2.forward(this.dropout.forward(
      this.relu.forward(this.fc1.forward(this.flatten.forward(input))),
    ));
  }
}

Training loop

The loop below follows the example's batch lifecycle. It explicitly reads the batch loss for logging, then runs backward and Adam before releasing temporary tensors.

Forward, backward, updateexamples/mnist/train.ts
const model = new SimpleMLP().to(device);
const parameters = model.parameters();
const criterion = new nn.CrossEntropyLoss();
const optimizer = nn.optim.Adam(parameters, { lr: 0.001 });

model.train();
let totalLoss = 0;
let totalSamples = 0;
for await (const batch of new DataLoader(trainData, 64, true)) {
  await optimizer.zeroGrad();
  const images = device.tensor(batch.images, {
    shape: [batch.size, 1, 28, 28],
  });
  const logits = model.forward(images);
  const loss = criterion.forward(logits, batch.labels);
  const batchLoss = (await loss.data())[0];
  if (batchLoss === undefined) throw new Error("CrossEntropyLoss returned no value.");
  await loss.backward();
  await optimizer.step();
  totalLoss += batchLoss * batch.size;
  totalSamples += batch.size;
  await loss.disposeGraph(parameters);
}

What happens on GPU

For WebGPU and WebGL2, TypeNN schedules gradient graph work and optimizer update kernels through a TypeShade Frame. Gradients, parameters, and Adam moment state are kept in Device-owned resident buffers during the update path. The CPU optimizer updates host arrays.

Explicit reads remain reads.Calling loss.data() or parameter.data() requests a host result. The sample loop above does not read the full parameter or gradient arrays back each step.

Dispose each batch graph

disposeGraph(model.parameters()) releases temporary tensors while retaining model parameters. Keep the Device alive for the training session and dispose it after the final batch and any pending work finish.

Open the full MNIST example ↗

Current scope

Tensor autograd covers selected operations. The current Module training example is the authoritative scope; this guide does not imply support for arbitrary user-defined kernels or all Tensor graph derivatives.