가이드

모델 학습

Module API에서 학습 루프는 명시적입니다. 기울기를 초기화하고 forward, 역전파, 파라미터 갱신을 실행한 뒤 배치 그래프를 해제합니다.

모델 정의

MNIST 예제는 ReLU와 Dropout을 포함한 784 → 128 → 10 MLP를 사용합니다. 아래 코드는 seed 가중치 옵션을 줄인 실제 모델 소스 구조입니다.

MNIST 모델examples/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))),
    ));
  }
}

학습 루프

아래 코드는 예제의 배치 수명주기를 따릅니다. 로깅할 loss 값을 명시적으로 읽은 다음 backward, Adam 갱신을 실행하고 임시 Tensor를 해제합니다.

Forward, backward, 갱신examples/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);
}

GPU 실행 경로

WebGPU와 WebGL2에서는 TypeNN이 기울기 그래프 연산과 옵티마이저 갱신 커널을 TypeShade Frame으로 실행합니다. 갱신 중 기울기, 파라미터, Adam 모멘트 상태는 Device 소유 Resident 버퍼에 유지됩니다. CPU 옵티마이저는 호스트 배열을 갱신합니다.

명시적으로 읽은 값은 호스트로 전달됩니다.loss.data()와 parameter.data()는 CPU 결과를 요청합니다. 위 루프는 매 단계 전체 파라미터나 기울기 배열을 읽지 않습니다.

각 배치 그래프 해제

disposeGraph(model.parameters())는 모델 파라미터를 보존하고 임시 Tensor를 해제합니다. 학습 세션이 끝나고 대기 중 작업을 마친 다음 Device를 해제하세요.

MNIST 전체 예제 보기 ↗

현재 지원 범위

Tensor 자동 미분은 선택된 연산을 지원합니다. Module 학습 예제를 현재 지원 범위로 보고, 임의 사용자 커널이나 모든 그래프 연산의 미분이 지원된다고 가정하지 마세요.