/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
with_cpu
apps/python/resnet/train.py
82 строки
3 KB
kolkir
Add tests for Linear layer and update matmul usage
28 мар 2025, 23:53
28 мар 2025, 23:53
bc788d8
Код
Авторство
О чём код?
import argparse import sys from tqdm import tqdm from pathlib import Path import time from resnet import ResNet18 from adept.data import CIFAR10, DataLoader from adept import device_t, dtype_t, Variable, no_grad, Tensor, TensorProperties, Shape from adept.optim import Adam from adept.loss import cross_entropy_with_logits from adept.serialize import FileInput, FileOutput device = device_t.CPU dtype = dtype_t.Float32 epochs = 5 batch_size = 8 lr = 0.01 num_classes = 10 def main(): parser = argparse.ArgumentParser( prog="ResNet trainer", description="Script trains ResNet with CIFAR-10 dataset" ) parser.add_argument("dataset_path", help="path to the CIFAR-10 root folder") parser.add_argument("-c", "--checkpoint", type=str, help="checkpoint file name") args = parser.parse_args(args=None if sys.argv[1:] else ["--help"]) train_dataset = CIFAR10(args.dataset_path, train=True) test_dataset = CIFAR10(args.dataset_path, train=False) test_dataloader = DataLoader(test_dataset, batch_size, True) train_dataloader = DataLoader(train_dataset, batch_size, True) resnet = ResNet18(num_classes) optimizer = Adam(resnet.parameters(), lr) if args.checkpoint and Path(args.checkpoint).exists(): input = FileInput(args.checkpoint) resnet.load(input) # TODO: load Adam params for epoch in tqdm(range(epochs), unit="epoch"): resnet.train() pbar = tqdm(train_dataloader, unit="batch") for b_i, batch in enumerate(pbar): x, y = batch out = resnet.forward(Variable(x, requires_grad=False)) loss = cross_entropy_with_logits(out, Variable(y, requires_grad=False)) if b_i % 64 == 0: pbar.set_postfix(loss=loss.data().float_at([0, 0])) loss.backward() optimizer.step() optimizer.zero_grad() # save checkpoint checkpoint_path = f"checkpoint_{epoch}_{int(time.time()*1000.0)}.pt" output = FileOutput(checkpoint_path) resnet.save(output) # TODO: save Adam params # test resnet.eval() with no_grad(): total_loss = Tensor.zero(TensorProperties(Shape([1]), device, dtype)) for b_i, batch in enumerate(train_dataloader): x, y = batch out = resnet.forward(Variable(x)) loss = cross_entropy_with_logits(out, Variable(y)) total_loss += loss.data() pbar.set_postfix(test_loss=total_loss.float_at([0]) / b_i) if __name__ == "__main__": main()