/
stupin
/
AutoDriver
Обзор
Документация
Войти
/
stupin
/
AutoDriver
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
Dev/Train.py
221 строка
8 KB
Stupin
init
25 июн 2025, 19:30
25 июн 2025, 19:30
a4828f6
Код
Авторство
О чём код?
from torch.cuda import is_available as cia from torchvision import transforms as T from torch.nn import functional as F #from PIL import Image, ImageDraw #from ultralytics import YOLO from model.ENet import ENet from torch import nn, optim from torchvision import io import argparse import torch import json import time import copy import os import warnings warnings.filterwarnings('ignore') class DLBO(): def __init__(self) -> None: self.epoch = 2 self.string = f'\033[97m{1:^5}\033[0m' self.is_train = True self.best_loss = [float('inf'), float('inf')] self.timer = time.perf_counter() print('\033[90mEpoch Train Loss Test Loss\033[0m') print(self.string, end='\r') def print(self, loss: float): ''' Prints loss and time to console ''' bl = self.best_loss[int(self.is_train)] if loss < bl: self.string += f' \033[92m↑ {str(loss)[:8]:0<8}\033[0m' else: self.string += f' \033[91m↓ {str(loss)[:8]:0<8}\033[0m' self.string += f' \033[90m({str(time.perf_counter() - self.timer)[:6]:0<6}s)\033[0m' if not self.is_train: print(self.string) self.string = f'\033[97m{self.epoch:^5}\033[0m' self.epoch += 1 self.best_loss[int(self.is_train)] = min(loss, bl) self.is_train = not self.is_train self.timer = time.perf_counter() print(self.string, end='\r') def end(self, reason: str = ''): ''' Closes console interaction and deletes itself ''' if reason: reason = 'Reason: \033[95m' + reason print('\033[90mTrain stopped.', reason, '\033[0m') del self class ArcheAgeDataset(): def __init__(self, path: str, split_coef: float, epoch: int, bpe: int, batch: int, device: torch.DeviceObjType = None): self.path = path self.device = device # Permutation prepare DATASET_SIZE = bpe * batch SPLIT = int(DATASET_SIZE * split_coef) SHUFFLE = torch.randperm(DATASET_SIZE, device=device) # Train permutation TRAIN_DRS = epoch * SPLIT * batch # Train Dataset Required Size train_imgs = SHUFFLE[:SPLIT] self.train_set = torch.clone(train_imgs) while self.train_set.shape[0] < TRAIN_DRS: shuffle = torch.randperm(train_imgs.shape[0], device=device) self.train_set = torch.cat((self.train_set, train_imgs[shuffle])) self.train_set = self.train_set[:TRAIN_DRS].reshape((epoch, train_imgs.shape[0], batch)) # Test permutation TEST_DRS = epoch * (DATASET_SIZE - SPLIT) * batch # Test Dataset Required Size test_imgs = SHUFFLE[SPLIT:] self.test_set = torch.clone(test_imgs) while self.test_set.shape[0] < TEST_DRS: shuffle = torch.randperm(test_imgs.shape[0], device=device) self.test_set = torch.cat((self.test_set, test_imgs[shuffle])) self.test_set = self.test_set[:TEST_DRS].reshape((epoch, test_imgs.shape[0], batch)) # Selecting startup set self.set = self.train_set def __len__(self): return self.set.shape[1] def __getitem__(self, idx): batch = torch.empty((0, 4, 360, 640), device=self.device) for i in self.set[idx]: batch = torch.cat((batch, io.read_image(f'{self.path}/{i}.png', mode=io.ImageReadMode.RGB_ALPHA).to(device=self.device).unsqueeze(0))) return batch[:, :3].float() / 256, 255 - batch[:, 3].long() def train(self): self.set = self.train_set def test(self): self.set = self.test_set #model = YOLO('weight/yolov8n-seg.pt') def main() -> None: # Parse input parameters parser = argparse.ArgumentParser() parser.add_argument('-i', '--input', type=str, default=None, help='Dataset directory path') parser.add_argument('-o', '--output', type=str, default='aaad.h5', help='Output model path') parser.add_argument('-n', '--model', type=str, default='ENet', help='Model names: ENet\nDefault: ENet') parser.add_argument('-s', '--size', type=str, default='640x360', help='Input images size\nDefault: 640x360') parser.add_argument('-c', '--split_coef', type=float, default=0.8, help='Train/Test split coefficient (0-1)\nDefalt: 0.8') parser.add_argument('-g', '--gpu', type=int, default=0, help='ID of GPU script should use') parser.add_argument('-b', '--batch', type=int, default=4, help='Batch size\nDefault: 4') parser.add_argument('-p', '--patience', type=int, default=5, help='Early stop patience (how much epochs without progress it will go through before ending train)\nDefault: 5') parser.add_argument('-e', '--epoch', type=int, default=20, help='Number of epochs\nDefault: 20') parser.add_argument('--bpe', type=int, default=100000, help='Batches per epoch\n--bpe * -b * -e should be more than dataset size so whole dataset will be used\nDefault: 100000') args = parser.parse_args() # Check and modify input parameters assert args.input, 'Specify dataset path with -i parameter' args.gpu = torch.device(f'cuda:{args.gpu}' if cia() else 'cpu') DATASET_SIZE = len(os.listdir(args.input)) - 1 if args.batch * args.bpe > DATASET_SIZE: args.bpe = DATASET_SIZE // args.batch # Open dataset dataset = ArcheAgeDataset(args.input, args.split_coef, args.epoch, args.bpe, args.batch, device=args.gpu) # Initialize model and training dependencies with open(f'{args.input}/classes.csv') as f: classes = len(f.read().split(',')) model = ENet(classes) model = model.to(args.gpu) criterion = F.cross_entropy optimizer = torch.optim.Adam(model.parameters(), lr=1e-1, weight_decay=0) lr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5, verbose=True) best_epoch_i = 0 best_test_loss = float('inf') best_model = copy.deepcopy(model) dlbo = DLBO() try: for epoch in range(args.epoch): # Train model.train() dataset.train() mean_train_loss = 0 for batch in range(len(dataset)): targets, labels = dataset[epoch, batch] pred = model(targets) loss = criterion(pred, labels) model.zero_grad() loss.backward() optimizer.step() mean_train_loss += float(loss) mean_train_loss /= len(dataset) dlbo.print(mean_train_loss) # Test model.eval() dataset.test() mean_test_loss = 0 with torch.no_grad(): for batch in range(len(dataset)): targets, labels = dataset[epoch, batch] pred = model(targets) loss = criterion(pred, labels) mean_test_loss += float(loss) mean_test_loss /= len(dataset) dlbo.print(mean_test_loss) # Epoch end actions lr_scheduler.step(mean_test_loss) if mean_test_loss < best_test_loss: best_epoch_i = epoch best_test_loss = mean_test_loss best_model = copy.deepcopy(model) elif epoch - best_epoch_i > args.patience: raise TimeoutError dlbo.end('Success') except KeyboardInterrupt: dlbo.end('Keyboard Interrupt') except TimeoutError: dlbo.end('Patience Interrupt') torch.save(best_model.state_dict(), args.output) if __name__ == '__main__': main()