가이드
모델 학습
Module API에서 학습 루프는 명시적입니다. 기울기를 초기화하고 forward, 역전파, 파라미터 갱신을 실행한 뒤 배치 그래프를 해제합니다.
모델 정의
MNIST 예제는 ReLU와 Dropout을 포함한 784 → 128 → 10 MLP를 사용합니다. 아래 코드는 seed 가중치 옵션을 줄인 실제 모델 소스 구조입니다.
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를 해제합니다.
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를 해제하세요.
현재 지원 범위
Tensor 자동 미분은 선택된 연산을 지원합니다. Module 학습 예제를 현재 지원 범위로 보고, 임의 사용자 커널이나 모든 그래프 연산의 미분이 지원된다고 가정하지 마세요.