Spaces:
Sleeping
Sleeping
| import gc | |
| import torch | |
| import numpy as np | |
| import pandas as pd | |
| from torch.nn.parallel import DistributedDataParallel | |
| from training.train import fit | |
| from model_zoo.models import define_model | |
| from data.dataset import CropDataset, ImageDataset, CoordsDataset | |
| from data.transforms import get_transfos | |
| from util.torch import seed_everything, count_parameters, save_model_weights | |
| def train(config, df_train, df_val, fold, log_folder=None, run=None): | |
| """ | |
| Train a crop model. | |
| Args: | |
| config (Config): Configuration parameters for training. | |
| df_train (pandas DataFrame): Metadata for training dataset. | |
| df_val (pandas DataFrame): Metadata for validation dataset. | |
| fold (int): Fold number for cross-validation. | |
| log_folder (str, optional): Folder for saving logs. Defaults to None. | |
| run: Neptune run. Defaults to None. | |
| Returns: | |
| tuple: A tuple containing predictions and metrics. | |
| """ | |
| if "crop" in config.pipe: | |
| dataset_class = CropDataset | |
| elif "coord" in config.pipe: | |
| dataset_class = CoordsDataset | |
| else: | |
| dataset_class = ImageDataset | |
| transfos = get_transfos( | |
| strength=config.aug_strength, | |
| resize=config.resize, | |
| crop=config.crop, | |
| use_keypoints="coords" in config.pipe, | |
| ) | |
| train_dataset = dataset_class( | |
| df_train, | |
| targets=config.targets, | |
| transforms=transfos, | |
| frames_chanel=config.frames_chanel, | |
| n_frames=config.n_frames, | |
| stride=config.stride, | |
| train=True, | |
| flip=config.flip if hasattr(config, "flip") else False, | |
| ) | |
| transfos = get_transfos( | |
| augment=False, | |
| resize=config.resize, | |
| crop=config.crop, | |
| use_keypoints="coords" in config.pipe, | |
| ) | |
| val_dataset = dataset_class( | |
| df_val, | |
| targets=config.targets, | |
| transforms=transfos, | |
| frames_chanel=config.frames_chanel, | |
| n_frames=config.n_frames, | |
| stride=config.stride, | |
| train=False, | |
| ) | |
| if config.pretrained_weights is not None: | |
| if any( | |
| [config.pretrained_weights.endswith(k) for k in [".pth", ".pt", ".bin"]] | |
| ): | |
| pretrained_weights = config.pretrained_weights | |
| else: # folder | |
| pretrained_weights = config.pretrained_weights + f"{config.name}_{fold}.pt" | |
| else: | |
| pretrained_weights = None | |
| model = define_model( | |
| config.name, | |
| drop_rate=config.drop_rate, | |
| drop_path_rate=config.drop_path_rate, | |
| pooling=config.pooling if hasattr(config, "pooling") else "avg", | |
| head_3d=config.head_3d, | |
| delta=config.delta if hasattr(config, "delta") else 2, | |
| n_frames=config.n_frames, | |
| num_classes=config.num_classes, | |
| num_classes_aux=config.num_classes_aux, | |
| n_channels=config.n_channels, | |
| pretrained_weights=pretrained_weights, | |
| reduce_stride=config.reduce_stride, | |
| verbose=(config.local_rank == 0), | |
| ).cuda() | |
| if config.distributed: | |
| model = DistributedDataParallel( | |
| model, | |
| device_ids=[config.local_rank], | |
| find_unused_parameters=False, | |
| broadcast_buffers=False, | |
| ) | |
| model.zero_grad(set_to_none=True) | |
| model.train() | |
| n_parameters = count_parameters(model) | |
| if config.local_rank == 0: | |
| print(f" -> {len(train_dataset)} training injuries") | |
| print(f" -> {len(val_dataset)} validation injuries") | |
| print(f" -> {n_parameters} trainable parameters\n") | |
| preds, metrics = fit( | |
| model, | |
| train_dataset, | |
| val_dataset, | |
| config.data_config, | |
| config.loss_config, | |
| config.optimizer_config, | |
| epochs=config.epochs, | |
| verbose_eval=config.verbose_eval, | |
| use_fp16=config.use_fp16, | |
| distributed=config.distributed, | |
| local_rank=config.local_rank, | |
| world_size=config.world_size, | |
| log_folder=log_folder, | |
| run=run, | |
| fold=fold, | |
| ) | |
| if (log_folder is not None) and (config.local_rank == 0): | |
| save_model_weights( | |
| model.module if config.distributed else model, | |
| f"{config.name}_{fold}.pt", | |
| cp_folder=log_folder, | |
| ) | |
| del (model, train_dataset, val_dataset) | |
| torch.cuda.empty_cache() | |
| gc.collect() | |
| return preds, metrics | |
| def k_fold(config, df, log_folder=None, run=None): | |
| """ | |
| Perform k-fold cross-validation training for a crop model. | |
| Args: | |
| config (dict): Configuration parameters for training. | |
| df (pandas DataFrame): Metadata. | |
| log_folder (str, optional): Folder for saving logs. Defaults to None. | |
| run: Neptune run. Defaults to None. | |
| """ | |
| folds = pd.read_csv(config.folds_file) | |
| df = df.merge(folds, how="left") | |
| df["fold"] = df["fold"].fillna(-1) | |
| all_metrics = [] | |
| for fold in range(config.k): | |
| if fold in config.selected_folds: | |
| if config.local_rank == 0: | |
| print( | |
| f"\n------------- Fold {fold + 1} / {config.k} -------------\n" | |
| ) | |
| seed_everything(config.seed + fold) | |
| df_train = df[df["fold"] != fold].reset_index(drop=True) | |
| df_val = df[df["fold"] == fold].reset_index(drop=True) | |
| if hasattr(config, "fix_train_crops"): | |
| if config.fix_train_crops: | |
| import re | |
| df_train["img_path"] = df_train["img_path"].apply( | |
| lambda x: re.sub( | |
| config.crop_folder, config.crop_folder[:-2] + "f/", x | |
| ) | |
| ) | |
| preds, metrics = train( | |
| config, | |
| df_train, | |
| df_val, | |
| fold, | |
| log_folder=log_folder, | |
| run=run, | |
| ) | |
| all_metrics.append(metrics) | |
| if log_folder is None: | |
| return | |
| if config.local_rank == 0: | |
| np.save(log_folder + f"pred_val_{fold}", preds) | |
| df_val.to_csv(log_folder + f"df_val_{fold}.csv", index=False) | |
| if config.local_rank == 0 and len(config.selected_folds): | |
| print("\n------------- CV Scores -------------\n") | |
| for k in all_metrics[0].keys(): | |
| avg = np.mean([m[k] for m in all_metrics]) | |
| print(f"- {k.split('_')[0][:7]} score\t: {avg:.3f}") | |
| if run is not None: | |
| run[f"global/{k}"] = avg | |
| if run is not None: | |
| run["global/logs"].upload(log_folder + "logs.txt") | |
| np.save(log_folder + f"pred_val_{fold}", preds) | |
| df_val.to_csv(log_folder + f"df_val_{fold}.csv", index=False) | |
| if config.fullfit and config.selected_folds != [0]: | |
| for ff in range(config.n_fullfit): | |
| if config.local_rank == 0: | |
| print( | |
| f"\n------------- Fullfit {ff + 1} / {config.n_fullfit} -------------\n" | |
| ) | |
| seed_everything(config.seed + ff) | |
| train( | |
| config, | |
| df, | |
| df[df["fold"] == 0].reset_index(drop=True), | |
| f"fullfit_{ff}", | |
| log_folder=log_folder, | |
| run=run, | |
| ) | |
| if run is not None: | |
| print() | |
| run.stop() | |