{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from timm.optim import create_optimizer_v2, optimizer_kwargs\nfrom timm.data.transforms import RandomResizedCropAndInterpolation\nfrom timm.loss import BinaryCrossEntropy\nfrom timm.data.transforms_factory import create_transform\nfrom timm.data import Mixup\nfrom timm.data.config import resolve_data_config\nimport timm\nimport argparse\nimport random\nfrom glob import glob\nimport yaml\nfrom timm.data.constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nfrom random import seed, shuffle\nfrom torchvision.datasets.folder import ImageFolder\nfrom typing import Any, Callable, cast, Dict, List, Optional, Tuple\nimport warnings\nfrom sklearn import metrics\nfrom tqdm import tqdm\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nimport torch.nn as nn\nimport torch\nimport cv2\nimport pandas as pd\nfrom sklearn.model_selection import KFold\nimport numpy as np\nimport pathlib\nimport os\nimport time\nfrom PIL import Image\nfrom torchvision.models import resnet50\nimport matplotlib.pyplot as plt\nfrom datetime import datetime\nimport transformers\nwarnings.filterwarnings('ignore')\n\n\ndef seed_everything(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nclass IMAGENET100(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        image_id = self.df.iloc[index]['image_id']\n        label = int(self.df.iloc[index]['class'])\n        path = os.path.join(self.img_dir, f'{image_id}.jpeg')\n        image = Image.open(path).convert('RGB')\n        if self.transform is not None:\n            image = self.transform(image)\n        return image, torch.tensor(label)\n\n\nclass IMAGENET100_test(Dataset):\n    def __init__(self, img_paths, transform=None):\n        self.img_paths = img_paths\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def __getitem__(self, index):\n        image_id = self.img_paths[index].split('/')[-1].split('.')[0]\n        path = self.img_paths[index]\n        image = Image.open(path).convert('RGB')\n        if self.transform is not None:\n            image = self.transform(image)\n        return image, image_id\n\n\ndef get_transform(image_size, mean, std):\n    return transforms.Compose([\n        transforms.Resize((image_size, image_size)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean, std)])\n\n\ndef read_yaml_string(yaml_string):\n    try:\n        data = yaml.safe_load(yaml_string)\n        return data\n    except yaml.YAMLError as e:\n        print(f\"Error while parsing YAML string: {e}\")\n\n\ndef convert_dict_to_argparse(dictionary):\n    parser = argparse.Namespace()\n    for key, value in dictionary.items():\n        setattr(parser, key, value)\n    return parser\n\n\ndef get_train_val_test_dataset(root_dir, image_size, mean, std):\n    train_df = pd.read_csv(os.path.join(root_dir, 'train.csv'))\n    val_df = pd.read_csv(os.path.join(root_dir, 'val.csv'))\n    test_files = glob(f'{root_dir}/test/*')\n    transform = get_transform(image_size, mean, std)\n    train_dataset = IMAGENET100(\n        train_df, os.path.join(root_dir, \"train\"), transform)\n    val_dataset = IMAGENET100(val_df, os.path.join(root_dir, \"val\"), transform)\n    test_dataset = IMAGENET100_test(test_files, transform)\n    return (train_dataset, val_dataset, test_dataset)\n\n\ndef create_loader(\n        dataset,\n        input_size,\n        batch_size,\n        is_training=False,\n        use_prefetcher=True,\n        no_aug=False,\n        re_prob=0.,\n        re_mode='const',\n        re_count=1,\n        re_split=False,\n        scale=None,\n        ratio=None,\n        hflip=0.5,\n        vflip=0.,\n        color_jitter=0.4,\n        auto_augment=None,\n        num_aug_splits=0,\n        interpolation='bilinear',\n        mean=IMAGENET_DEFAULT_MEAN,\n        std=IMAGENET_DEFAULT_STD,\n        num_workers=1,\n        crop_pct=None,\n        tf_preprocessing=False,\n        **kwargs\n):\n    re_num_splits = 0\n    if re_split:\n        # apply RE to second half of batch if no aug split otherwise line up with aug split\n        re_num_splits = num_aug_splits or 2\n    dataset.transform = create_transform(\n        input_size,\n        is_training=is_training,\n        use_prefetcher=use_prefetcher,\n        no_aug=no_aug,\n        scale=scale,\n        ratio=ratio,\n        hflip=hflip,\n        vflip=vflip,\n        color_jitter=color_jitter,\n        auto_augment=auto_augment,\n        interpolation=interpolation,\n        mean=mean,\n        std=std,\n        crop_pct=crop_pct,\n        tf_preprocessing=tf_preprocessing,\n        re_prob=re_prob,\n        re_mode=re_mode,\n        re_count=re_count,\n        re_num_splits=re_num_splits,\n        separate=num_aug_splits > 0,\n    )\n    print(dataset.transform)\n    loader = None\n    if kwargs['data_type'] == 'train':\n        loader = DataLoader(dataset, batch_size=batch_size,\n                            shuffle=is_training, num_workers=num_workers,drop_last=True)\n    else:\n        loader = DataLoader(dataset, batch_size=batch_size,\n                            shuffle=is_training, num_workers=num_workers)\n\n    return loader\n\n\nclass Backbone(nn.Module):\n    def __init__(self, args):\n        super(Backbone, self).__init__()\n        self.encoder = timm.create_model(\n            args.model,\n            pretrained=args.pretrained,\n            in_chans=3,\n            num_classes=args.num_classes,\n            drop_rate=args.drop,\n            drop_path_rate=args.drop_path,\n            drop_block_rate=args.drop_block,\n            global_pool=args.gp,\n            bn_momentum=args.bn_momentum,\n            bn_eps=args.bn_eps,\n            scriptable=args.torchscript,\n            checkpoint_path=args.initial_checkpoint,\n        )\n\n    def forward(self, x):\n        return self.encoder(x)\n\n\nclass Trainer():\n    def __init__(self,\n                 train_loader,\n                 val_loader,\n                 test_loader,\n                 device,\n                 model,\n                 optimizer,\n                 scheduler,\n                 epochs,\n                 model_save_dir,\n                 train_criterion,\n                 val_criterion,\n                 mixup_fn,\n                 test_output_dir,\n                 train_time):\n\n        self.epochs = epochs\n        self.epoch = 0\n        self.accelarator = device\n        self.model = model\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n\n        self.test_output_dir = test_output_dir\n        self.model_save_dir = model_save_dir\n\n        self.train_loader = train_loader\n        self.validation_loader = val_loader\n        self.test_loader = test_loader\n\n        self.mixup_fn = mixup_fn\n        self.train_criterion = train_criterion\n        self.val_criterion = val_criterion\n        self.train_time = train_time\n        self.train_start = -1\n        self.train_end = -1\n\n    def get_metrics(self, predictions, actual, isTensor=False):\n        if isTensor:\n            p = predictions.detach().cpu().numpy()\n            a = actual.detach().cpu().numpy()\n        else:\n            p = predictions\n            a = actual\n        accuracy = metrics.accuracy_score(y_pred=p, y_true=a)\n        return {\n            \"accuracy\": accuracy\n        }\n\n    def get_lr(self, optimizer):\n        for param_group in optimizer.param_groups:\n            return param_group['lr']\n\n    def train_step(self):\n        self.model.train()\n        print(\"Train Loop!\")\n        running_loss_train = 0.0\n        num_train = 0\n        train_predictions = np.array([])\n        train_labels = np.array([])\n\n        for images, labels in tqdm(self.train_loader):\n            images = images.to(self.accelarator)\n            labels = labels.to(self.accelarator)\n            images, labels = self.mixup_fn(images, labels)\n\n            num_train += labels.shape[0]\n            self.optimizer.zero_grad()\n            outputs = self.model(images)\n            loss = self.train_criterion(outputs, labels)\n            running_loss_train += loss.item()\n            loss.backward()\n            self.optimizer.step()\n            self.scheduler.step()\n            self.lr = self.get_lr(self.optimizer)\n\n            if time.time() >= self.train_end:\n                return TRAIN_END_FLAG\n\n        print(f\"Train Loss: {running_loss_train/num_train}\")\n\n        return {\n            'loss': running_loss_train/num_train,\n        }\n\n    def val_step(self):\n        val_predictions = np.array([])\n        val_labels = np.array([])\n        running_loss_val = 0.0\n        num_val = 0\n        self.model.eval()\n        with torch.no_grad():\n            print(\"Validation Loop!\")\n            for images, labels in tqdm(self.validation_loader):\n                images = images.to(self.accelarator)\n                labels = labels.to(self.accelarator)\n                outputs = self.model(images)\n                num_val += labels.shape[0]\n                _, preds = torch.max(outputs, 1)\n                val_predictions = np.concatenate(\n                    (val_predictions, preds.detach().cpu().numpy()))\n                val_labels = np.concatenate(\n                    (val_labels, labels.detach().cpu().numpy()))\n\n                loss = self.val_criterion(outputs, labels)\n                running_loss_val += loss.item()\n                if time.time() >= self.train_end:\n                    return TRAIN_END_FLAG\n            val_metrics = self.get_metrics(val_predictions, val_labels)\n            print(f\"Validation Loss: {running_loss_val/num_val}\")\n            print(f\"Val Accuracy Metric: {val_metrics['accuracy']} \")\n            return {\n                'loss': running_loss_val/num_val,\n                'accuracy': val_metrics['accuracy'],\n            }\n\n    def test_step(self):\n        test_image_ids = np.array([])\n        test_predictions = np.array([])\n        best_model_path = os.path.join(self.model_save_dir,'best_val_accuracy.pt')\n        if os.path.exists(best_model_path):\n            checkpoint = torch.load(\n                f\"{self.model_save_dir}/best_val_accuracy.pt\")\n            start_epoch = checkpoint['epoch']\n            print(\n                f\"Model already trained for {start_epoch} epochs.\")\n            print(self.model.load_state_dict(checkpoint['model1_weights']))\n        self.model.eval()\n        with torch.no_grad():\n            print(\"Test Loop!\")\n            for images, image_ids in tqdm(self.test_loader):\n                images = images.to(self.accelarator)\n                outputs = self.model(images)\n                _, preds = torch.max(outputs, 1)\n                test_image_ids = np.concatenate((test_image_ids, image_ids))\n                test_predictions = np.concatenate(\n                    (test_predictions, preds.detach().cpu().numpy()))\n\n            pd.DataFrame({\n                'image_id': test_image_ids,\n                'label': test_predictions.astype(int)\n            }).to_csv(os.path.join(self.test_output_dir, 'submission.csv'), index=False)\n\n    def run(self, run_test=True):\n        best_validation_loss = float('inf')\n        best_validation_accuracy = 0\n        self.train_start = time.time()\n        self.train_end = self.train_start + self.train_time\n        for epoch in range(self.epochs):\n            print(\"=\"*31)\n            print(f\"{'-'*10} Epoch {epoch+1}/{self.epochs} {'-'*10}\")\n            train_logs = self.train_step()\n            if train_logs == TRAIN_END_FLAG:\n                torch.save({\n                    'model1_weights': self.model.state_dict(),\n                    'optimizer_state': self.optimizer.state_dict(),\n                    'scheduler_state': self.scheduler.state_dict(),\n                    'epoch': epoch+1,\n                }, f\"{self.model_save_dir}/last.pt\")\n                break\n            val_logs = self.val_step()\n            if val_logs == TRAIN_END_FLAG:\n                torch.save({\n                    'model1_weights': self.model.state_dict(),\n                    'optimizer_state': self.optimizer.state_dict(),\n                    'scheduler_state': self.scheduler.state_dict(),\n                    'epoch': epoch+1,\n                }, f\"{self.model_save_dir}/last.pt\")\n                break\n            self.epoch = epoch\n            if val_logs[\"loss\"] < best_validation_loss:\n                best_validation_loss = val_logs[\"loss\"]\n                torch.save({\n                    'model1_weights': self.model.state_dict(),\n                    'optimizer_state': self.optimizer.state_dict(),\n                    'scheduler_state': self.scheduler.state_dict(),\n                    'epoch': epoch+1,\n                }, f\"{self.model_save_dir}/best_val_loss.pt\")\n            if val_logs['accuracy'] > best_validation_accuracy:\n                best_validation_accuracy = val_logs['accuracy']\n                torch.save({\n                    'model1_weights': self.model.state_dict(),\n                    'optimizer_state': self.optimizer.state_dict(),\n                    'scheduler_state': self.scheduler.state_dict(),\n                    'epoch': epoch+1,\n                }, f\"{self.model_save_dir}/best_val_accuracy.pt\")\n        if run_test:\n            self.test_step()\n\n        return {\n            'best_accuracy': best_validation_accuracy,\n            'best_loss': best_validation_loss,\n        }\n\n\nif __name__ == \"__main__\":\n    yaml_string = '''\n        aa: rand-m6-mstd0.5-inc1\n        distributed: false\n        device: 0\n        rank: 0\n        amp: true\n        apex_amp: false\n        aug_repeats: 0\n        aug_splits: 0\n        batch_size: 1024\n        bce_loss: true\n        bn_eps: null\n        bn_momentum: null\n        bn_tf: false\n        channels_last: true\n        checkpoint_hist: 10\n        clip_grad: null\n        clip_mode: norm\n        color_jitter: 0.4\n        cooldown_epochs: 10\n        crop_pct: 0.95\n        cutmix: 1.0\n        cutmix_minmax: null\n        data_dir: /imagenet\n        decay_epochs: 100\n        decay_rate: 0.1\n        dist_bn: reduce\n        drop: 0.0\n        drop_block: null\n        drop_connect: null\n        drop_path: null\n        epoch_repeats: 0.0\n        epochs: 100\n        eval_metric: top1\n        experiment: \"\"\n        gp: null\n        hflip: 0.5\n        img_size: 160\n        initial_checkpoint: \"\"\n        input_size: null\n        interpolation: \"\"\n        jsd_loss: false\n        local_rank: 0\n        log_interval: 50\n        log_wandb: false\n        lr: 0.008\n        lr_cycle_decay: 0.5\n        lr_cycle_limit: 1\n        lr_cycle_mul: 1.0\n        lr_k_decay: 1.0\n        lr_noise: null\n        lr_noise_pct: 0.67\n        lr_noise_std: 1.0\n        mean: null\n        min_lr: 1.0e-06\n        mixup: 0.1\n        mixup_mode: batch\n        mixup_off_epoch: 0\n        mixup_prob: 1.0\n        mixup_switch_prob: 0.5\n        model: resnet50\n        model_ema: false\n        model_ema_decay: 0.9998\n        model_ema_force_cpu: false\n        momentum: 0.9\n        native_amp: false\n        no_aug: false\n        no_prefetcher: true\n        no_resume_opt: false\n        num_classes: null\n        opt: lamb\n        opt_betas: null\n        opt_eps: null\n        output: \"\"\n        patience_epochs: 10\n        pin_mem: false\n        pretrained: false\n        ratio:\n        - 0.75\n        - 1.3333333333333333\n        recount: 1\n        recovery_interval: 0\n        remode: pixel\n        reprob: 0.0\n        resplit: false\n        resume: \"\"\n        save_images: false\n        scale:\n        - 0.08\n        - 1.0\n        sched: cosine\n        seed: 0\n        smoothing: 0.0\n        split_bn: false\n        start_epoch: null\n        std: null\n        sync_bn: false\n        torchscript: false\n        train_interpolation: random\n        train_split: train\n        tta: 0\n        use_multi_epochs_loader: false\n        val_split: validation\n        validation_batch_size: null\n        vflip: 0.0\n        warmup_epochs: 5\n        warmup_lr: 0.0001\n        weight_decay: 0.02\n        workers: 4\n        world_size: 1\n        bce_target_thresh: 0.2\n    '''\n\n    arguments = read_yaml_string(yaml_string)\n    arg_parser = convert_dict_to_argparse(arguments)\n\n    train_parser = argparse.ArgumentParser(\n        description=\"Arguments for training baseline on ImageNet100\")\n    train_parser.add_argument('--config', default='./config.yaml', type=str, metavar='FILE',\n                              help='YAML config file specifying default arguments')\n    train_parser.add_argument('--root_dir', default='./', type=str)\n    train_parser.add_argument('--epochs', default=100, type=int)\n    train_parser.add_argument('--batch_size', default=80, type=int)\n    train_parser.add_argument('--image_size', default=160, type=int)\n    train_parser.add_argument('--seed', default=42, type=int)\n    train_parser.add_argument('--num_classes', default=100, type=int)\n    train_parser.add_argument('--num_workers', default=2, type=int)\n    train_parser.add_argument('--lr', default=1e-3, type=float)\n    train_parser.add_argument('--output_dir', default='./', type=str)\n    train_parser.add_argument('--train_time', default=8.8*60*60, type=int)\n    parser = argparse.ArgumentParser(\n        description=\"Complete parser\")\n    args1 = train_parser.parse_args()\n    args2 = arg_parser\n    args = argparse.Namespace()\n    for key, value in vars(args2).items():\n        setattr(args, key, value)\n    for key, value in vars(args1).items():\n        setattr(args, key, value)\n    args_text = yaml.safe_dump(args.__dict__, default_flow_style=False)\n    args.prefetcher = False\n    \n    \n    tup = torch.cuda.mem_get_info()\n    torch.cuda.set_per_process_memory_fraction((6*(1024)**3)/tup[1], 0)\n\n    EPOCHS = args.epochs\n    ACCELARATOR = f'cuda' if torch.cuda.is_available() else 'cpu'\n\n    IMAGE_SIZE = args.img_size\n    BATCH_SIZE = args.batch_size\n\n    TRAIN_END_FLAG = -1\n\n    NUM_CLASSES = args.num_classes\n    SEED = args.seed\n    WARMUP_EPOCHS = args.warmup_epochs\n    NUM_WORKERS = args.num_workers\n    ROOT_DIR = args.root_dir\n    MEAN = IMAGENET_DEFAULT_MEAN\n    STD = IMAGENET_DEFAULT_STD\n    DECAY_FACTOR = 1\n    OUTPUT_DIR = args.output_dir\n    MODEL_SAVE_DIR = OUTPUT_DIR\n\n    seed_everything(SEED)\n\n    model1 = Backbone(args)\n    model1.to(ACCELARATOR)\n    for param in model1.parameters():\n        param.requires_grad = True\n    data_config = resolve_data_config(\n        vars(args), model=model1.encoder, verbose=True)\n\n    print(f\"Baseline model:\")\n\n    train_loss_fn = BinaryCrossEntropy(\n        target_threshold=args.bce_target_thresh).to()\n    train_loss_fn = train_loss_fn.to(device=ACCELARATOR)\n    validate_loss_fn = nn.CrossEntropyLoss().to(device=ACCELARATOR)\n\n    train_dataset, val_dataset, test_dataset = get_train_val_test_dataset(\n        args.root_dir, image_size=IMAGE_SIZE, mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD)\n    train_loader = create_loader(\n        train_dataset,\n        input_size=data_config['input_size'],\n        batch_size=args.batch_size,\n        is_training=True,\n        use_prefetcher=args.prefetcher,\n        no_aug=args.no_aug,\n        re_prob=args.reprob,\n        re_mode=args.remode,\n        re_count=args.recount,\n        re_split=args.resplit,\n        scale=args.scale,\n        ratio=args.ratio,\n        hflip=args.hflip,\n        vflip=args.vflip,\n        color_jitter=args.color_jitter,\n        auto_augment=args.aa,\n        num_aug_repeats=args.aug_repeats,\n        num_aug_splits=0,\n        interpolation=args.train_interpolation,\n        mean=data_config['mean'],\n        std=data_config['std'],\n        num_workers=args.workers,\n        distributed=args.distributed,\n        collate_fn=None,\n        pin_memory=args.pin_mem,\n        use_multi_epochs_loader=args.use_multi_epochs_loader,\n        data_type='train'\n    )\n    validation_loader = create_loader(\n        val_dataset,\n        input_size=data_config['input_size'],\n        batch_size=args.validation_batch_size or args.batch_size,\n        is_training=False,\n        use_prefetcher=args.prefetcher,\n        interpolation=data_config['interpolation'],\n        mean=data_config['mean'],\n        std=data_config['std'],\n        num_workers=args.num_workers,\n        distributed=args.distributed,\n        crop_pct=data_config['crop_pct'],\n        pin_memory=args.pin_mem,\n        data_type='val'\n    )\n    test_loader = create_loader(\n        test_dataset,\n        input_size=data_config['input_size'],\n        batch_size=args.validation_batch_size or args.batch_size,\n        is_training=False,\n        use_prefetcher=args.prefetcher,\n        interpolation=data_config['interpolation'],\n        mean=data_config['mean'],\n        std=data_config['std'],\n        num_workers=args.num_workers,\n        distributed=args.distributed,\n        crop_pct=data_config['crop_pct'],\n        pin_memory=args.pin_mem,\n        data_type='test'\n    )\n\n    print(\n        f\"Length of train loader: {len(train_loader)},Validation loader: {(len(validation_loader))}, Test Loader:{(len(test_loader))}\")\n\n    collate_fn = None\n    mixup_fn = None\n    mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None\n    if mixup_active:\n        mixup_args = dict(\n            mixup_alpha=args.mixup,\n            cutmix_alpha=args.cutmix,\n            cutmix_minmax=args.cutmix_minmax,\n            prob=args.mixup_prob,\n            switch_prob=args.mixup_switch_prob,\n            mode=args.mixup_mode,\n            label_smoothing=args.smoothing,\n            num_classes=args.num_classes\n        )\n    mixup_fn = Mixup(**mixup_args)\n    steps_per_epoch = len(train_dataset)//(BATCH_SIZE)\n    if len(train_dataset) % BATCH_SIZE != 0:\n        steps_per_epoch += 1\n    optimizer = torch.optim.Adam(model1.parameters(), lr=args.lr)\n\n    scheduler = transformers.get_cosine_schedule_with_warmup(\n        optimizer, WARMUP_EPOCHS*steps_per_epoch, DECAY_FACTOR*EPOCHS*steps_per_epoch)\n\n    trainer = Trainer(\n        train_loader=train_loader,\n        val_loader=validation_loader,\n        test_loader=test_loader,\n        device=ACCELARATOR,\n        model=model1,\n        optimizer=optimizer,\n        scheduler=scheduler,\n        epochs=EPOCHS,\n        model_save_dir=MODEL_SAVE_DIR,\n        train_criterion=train_loss_fn,\n        val_criterion=validate_loss_fn,\n        mixup_fn=mixup_fn,\n        test_output_dir=OUTPUT_DIR,\n        train_time=args.train_time)\n    trainer.run(True)","metadata":{},"execution_count":null,"outputs":[]}]}