Download main.py from ODELIA-AI/ABMIL: direct link, hf CLI and curl.
- Browser
- Download file 29.7 kB
-
https://huggingface.co/ODELIA-AI/ABMIL/resolve/main/main.py
- Command line
-
hf download hf://ODELIA-AI/ABMIL/main.py
-
curl -L -o main.py https://huggingface.co/ODELIA-AI/ABMIL/resolve/main/main.py
29.7 kB
| import os | |
| import torch | |
| import wandb | |
| from dataset import BreastMRI_ABMIL_NII, MammographyDataset, BreastMRI, BreastMRI_ABMIL | |
| #from src.models_all.models import create_soat_cnn_models | |
| import torch.optim as optim | |
| import pandas as pd | |
| import argparse | |
| import torch.nn as nn | |
| from pathlib import Path | |
| from torch.utils.data import WeightedRandomSampler | |
| from torchinfo import summary | |
| from src.trainer_tester.supervised_trainer import FullySupervisedMultiViewTrainer, MultiInstanceTrainer | |
| from src.utils.utils_functions import compute_sample_weights, log_scalar_values, set_seed | |
| from src.data_loader.mammo_transforms import BasicTransforms, TrainTransform, TrainTransformBaseline | |
| from src.utils.plotter import display_confusion_matrix | |
| from sklearn.metrics import classification_report | |
| import numpy as np | |
| import pandas as pd | |
| from pathlib import Path | |
| from src.utils.dataset_utils import extract_target_names | |
| #import utils | |
| import torch | |
| import torch.nn as nn | |
| import timm | |
| from focal_loss.focal_loss import FocalLoss | |
| class MultiViewSwinModel(nn.Module): | |
| def __init__(self, model_name="swin_tiny_patch4_window7_224.ms_in22k", | |
| pretrained=True, num_classes=3, attn_dim=1024, num_heads=4): | |
| super(MultiViewSwinModel, self).__init__() | |
| # Load backbone without classification head | |
| self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0) | |
| self.feature_dim = self.backbone.num_features # should be 1024 for swin_base | |
| # Self-attention across views (3 views → sequence length 3) | |
| self.attn = nn.MultiheadAttention(embed_dim=self.feature_dim, num_heads=num_heads, batch_first=True) | |
| # Final classifier | |
| self.classifier = nn.Linear(self.feature_dim, num_classes) | |
| def extract_feat(self, x): | |
| # Get unpooled features (B, H*W, C) and then pooled vector (B, C) | |
| features = self.backbone.forward_features(x) # (B, H*W, C) | |
| pooled = self.backbone.forward_head(features, pre_logits=True) # (B, C) | |
| return pooled | |
| def forward(self, img1, img2, img3): | |
| # Per-view embeddings | |
| f1 = self.extract_feat(img1) # [B, C] | |
| f2 = self.extract_feat(img2) | |
| f3 = self.extract_feat(img3) | |
| # Stack views to form a sequence [B, 3, C] | |
| views = torch.stack([f1, f2, f3], dim=1) | |
| # Self-attention across views | |
| attn_out, _ = self.attn(views, views, views) # [B, 3, C] | |
| # Aggregate (mean pooling across 3 views) | |
| fused = attn_out.mean(dim=1) # [B, C] | |
| return self.classifier(fused) | |
| class MultiViewSwinCrossAttn(nn.Module): | |
| """ | |
| Cross‑attention Swin for 3‑view breast‑MRI (pre / post / subtraction). | |
| """ | |
| def __init__( | |
| self, | |
| model_name: str = "swin_tiny_patch4_window7_224.ms_in22k", | |
| pretrained: bool = True, | |
| num_classes: int = 3, | |
| num_xattn_layers: int = 2, | |
| num_heads: int = 4, | |
| dropout: float = 0.1, | |
| ): | |
| super().__init__() | |
| # 1) Swin backbone WITHOUT final linear head | |
| self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0) | |
| self.embed_dim = self.backbone.num_features # 1024 for swin‑base | |
| # 2) Learnable CLS token (like ViT) and view‑type embeddings | |
| self.cls_token = nn.Parameter(torch.zeros(1, 1, self.embed_dim)) | |
| self.view_embed = nn.Parameter(torch.zeros(3, 1, self.embed_dim)) # pre/post/sub | |
| # 3) Cross‑attention encoder (TransformerEncoder) | |
| encoder_layer = nn.TransformerEncoderLayer( | |
| d_model=self.embed_dim, | |
| nhead=num_heads, | |
| dim_feedforward=self.embed_dim * 4, | |
| dropout=dropout, | |
| activation="gelu", | |
| batch_first=True, | |
| norm_first=True, | |
| ) | |
| self.xattn = nn.TransformerEncoder(encoder_layer, num_layers=num_xattn_layers) | |
| # 4) Classification head | |
| self.norm = nn.LayerNorm(self.embed_dim) | |
| self.classifier = nn.Linear(self.embed_dim, num_classes) | |
| # init | |
| nn.init.trunc_normal_(self.cls_token, std=0.02) | |
| nn.init.trunc_normal_(self.view_embed, std=0.02) | |
| # ------------------------------------------------------------------ # | |
| def _img_embed(self, x: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Forward a single image through Swin and return pooled vector (B, C). | |
| """ | |
| feats = self.backbone.forward_features(x) # (B, H*W, C) | |
| pooled = self.backbone.forward_head(feats, pre_logits=True) # (B, C) | |
| return pooled | |
| def forward(self, pre, post, sub): | |
| """ | |
| Args: | |
| pre, post, sub: tensors (B, 3, 224, 224) | |
| Returns: | |
| logits (B, num_classes) | |
| """ | |
| # 1) Per‑view embeddings | |
| z_pre = self._img_embed(pre) # (B, C) | |
| z_post = self._img_embed(post) | |
| z_sub = self._img_embed(sub) | |
| # 2) Assemble sequence: [CLS] + three view tokens | |
| B = z_pre.size(0) | |
| cls = self.cls_token.expand(B, -1, -1) # (B, 1, C) | |
| views = torch.stack([z_pre, z_post, z_sub], dim=1) # (B, 3, C) | |
| # add view‑type embeddings (broadcast over batch) | |
| views = views + self.view_embed.transpose(0, 1) # (B, 3, C) | |
| tokens = torch.cat([cls, views], dim=1) # (B, 4, C) | |
| # 3) Cross‑attention | |
| tokens = self.xattn(tokens) # (B, 4, C) | |
| # 4) CLS pooling → head | |
| out = self.classifier(self.norm(tokens[:, 0])) # (B, num_classes) | |
| return out | |
| import torch | |
| import torch.nn as nn | |
| import timm | |
| import torch | |
| import torch.nn as nn | |
| import timm | |
| import math | |
| import numpy as np | |
| class CrossModalAttentionABMIL_Swin(nn.Module): | |
| """ | |
| Swin‑based Multiple‑Instance model with per‑slice cross‑modal attention | |
| over (pre, post, sub) features, followed by ABMIL pooling. | |
| Input : (B, 32, 3, 224, 224) # 3 modalities stacked in channel dim | |
| Output : logits (B, num_classes), attention‑per‑slice (B, 32) | |
| """ | |
| def __init__( | |
| self, | |
| model_name: str = "swin_tiny_patch4_window7_224.ms_in22k", | |
| pretrained: bool = True, | |
| num_classes: int = 3, | |
| hidden_dim: int = 256, | |
| cross_attn_heads: int = 4, | |
| dropout: float = 0.1, | |
| ): | |
| super().__init__() | |
| # ------------------------------------------------------------ | |
| # 1) Shared Swin backbone (no classifier head) | |
| # ------------------------------------------------------------ | |
| self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0) | |
| self.embed_dim = self.backbone.num_features # 768 for swin_tiny | |
| # ------------------------------------------------------------ | |
| # 2) Per‑slice cross‑modal fusion | |
| # Input: three embeddings (pre, post, sub) -> fused embedding | |
| # ------------------------------------------------------------ | |
| encoder_layer = nn.TransformerEncoderLayer( | |
| d_model=self.embed_dim, | |
| nhead=cross_attn_heads, | |
| dim_feedforward=self.embed_dim * 4, | |
| dropout=dropout, | |
| activation="gelu", | |
| batch_first=True, | |
| norm_first=True, | |
| ) | |
| # One small transformer (2 layers) that attends across 3 tokens | |
| self.slice_fuser = nn.TransformerEncoder(encoder_layer, num_layers=2) | |
| # Token type embedding to differentiate modalities | |
| self.mod_embed = nn.Parameter(torch.randn(3, 1, self.embed_dim)) | |
| # ------------------------------------------------------------ | |
| # 3) ABMIL Attention across 32 fused slices | |
| # ------------------------------------------------------------ | |
| self.attn_V = nn.Linear(self.embed_dim, hidden_dim) | |
| self.attn_U = nn.Linear(hidden_dim, 1) | |
| # ------------------------------------------------------------ | |
| # 4) Classifier head | |
| # ------------------------------------------------------------ | |
| self.norm = nn.LayerNorm(self.embed_dim) | |
| self.dropout = nn.Dropout(dropout) | |
| self.classifier = nn.Linear(self.embed_dim, num_classes) | |
| # Init | |
| nn.init.trunc_normal_(self.mod_embed, std=0.02) | |
| nn.init.trunc_normal_(self.attn_V.weight, std=0.02) | |
| nn.init.trunc_normal_(self.attn_U.weight, std=0.02) | |
| # ------------------------------------------------------------ | |
| def _encode_slices(self, x: torch.Tensor) -> torch.Tensor: | |
| """ | |
| x : (B*N, 1, 224, 224) -> Swin pooled feats (B*N, embed_dim) | |
| """ | |
| feats = self.backbone.forward_features(x) | |
| return self.backbone.forward_head(feats, pre_logits=True) | |
| # ------------------------------------------------------------ | |
| def forward(self, volume: torch.Tensor): | |
| """ | |
| volume : (B, 32, 3, 224, 224) | |
| """ | |
| B, N, C_img, H, W = volume.shape # C_img = 3 modalities | |
| assert C_img == 3, "Expect channel dim = 3 (pre, post, sub)" | |
| # -------------------------------------------------------- | |
| # 1) Split modalities, flatten, and encode with Swin | |
| # -------------------------------------------------------- | |
| # Each is (B*N, 1, H, W) | |
| def expand3(x): return x.expand(-1, 3, -1, -1) | |
| pre = expand3(volume[:, :, 0, :, :].contiguous().view(B * N, 1, H, W)) | |
| post = expand3(volume[:, :, 1, :, :].contiguous().view(B * N, 1, H, W)) | |
| sub = expand3(volume[:, :, 2, :, :].contiguous().view(B * N, 1, H, W)) | |
| feat_pre = self._encode_slices(pre).view(B, N, -1) # (B, N, C) | |
| feat_post = self._encode_slices(post).view(B, N, -1) | |
| feat_sub = self._encode_slices(sub).view(B, N, -1) | |
| # -------------------------------------------------------- | |
| # 2) Cross‑modal attention fusion (per slice) | |
| # -------------------------------------------------------- | |
| # Build token sequence [pre, post, sub] for each slice | |
| # Shape before fuser: (B, N, 3, C) -> we fuse along dim=2 | |
| slice_tokens = torch.stack([feat_pre, feat_post, feat_sub], dim=2) | |
| # Add modality embeddings | |
| slice_tokens = slice_tokens + self.mod_embed.transpose(0, 1) # (B, N, 3, C) | |
| slice_tokens = slice_tokens.view(B * N, 3, self.embed_dim) # (B*N, 3, C) | |
| # Transformer encoder attends across the 3 tokens | |
| fused = self.slice_fuser(slice_tokens)[:, 0] # take CLS‑like first token | |
| fused = fused.view(B, N, self.embed_dim) # (B, 32, C) | |
| # -------------------------------------------------------- | |
| # 3) ABMIL attention over 32 fused slices | |
| # -------------------------------------------------------- | |
| A = torch.tanh(self.attn_V(fused)) # (B, N, hidden) | |
| A = self.attn_U(A) # (B, N, 1) | |
| A = torch.softmax(A, dim=1) # (B, N, 1) | |
| patient_feat = (A * fused).sum(dim=1) # (B, C) | |
| # -------------------------------------------------------- | |
| # 4) Head | |
| # -------------------------------------------------------- | |
| logits = self.classifier(self.dropout(self.norm(patient_feat))) | |
| return logits, A.squeeze(-1) # (B, num_classes), (B, 32) | |
| class ABMIL_Swin(nn.Module): | |
| """ | |
| Attention‑based Multiple‑Instance Learning (ABMIL) model | |
| for 3‑channel breast MRI slice triplets (pre / post / sub). | |
| Input shape : (B, 32, 3, 224, 224) | |
| Output : logits (B, num_classes) + attention weights (B, 32) | |
| """ | |
| def __init__( | |
| self, | |
| model_name: str = "swin_tiny_patch4_window7_224.ms_in22k", | |
| pretrained: bool = True, | |
| num_classes: int = 2, | |
| hidden_dim: int = 256, | |
| dropout: float = 0.1, | |
| ): | |
| super().__init__() | |
| # 1) Swin backbone WITHOUT final linear head | |
| self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0) | |
| self.embed_dim = self.backbone.num_features # e.g. 768 (tiny), 1024 (base) | |
| # 2) Attention network (ABMIL) | |
| self.attn_V = nn.Linear(self.embed_dim, hidden_dim) | |
| self.attn_U = nn.Linear(hidden_dim, 1) | |
| # 3) Patient‑level classifier | |
| self.norm = nn.LayerNorm(self.embed_dim) | |
| self.dropout = nn.Dropout(dropout) | |
| self.classifier = nn.Linear(self.embed_dim, num_classes) | |
| # Optional init | |
| nn.init.trunc_normal_(self.attn_V.weight, std=0.02) | |
| nn.init.trunc_normal_(self.attn_U.weight, std=0.02) | |
| # ------------------------------------------------------------------ # | |
| def _slice_embed(self, x: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Forward one or more images through Swin and return pooled features. | |
| Args: | |
| x : (N, 3, 224, 224) | |
| Returns: | |
| pooled : (N, C) | |
| """ | |
| feats = self.backbone.forward_features(x) # (N, H*W, C) | |
| pooled = self.backbone.forward_head(feats, pre_logits=True) # (N, C) | |
| return pooled | |
| # ------------------------------------------------------------------ # | |
| def forward(self, volume: torch.Tensor): | |
| """ | |
| Args: | |
| volume : (B, 32, 3, 224, 224) | |
| Returns: | |
| logits : (B, num_classes) | |
| attn_w : (B, 32) -- attention per slice (sums to 1) | |
| """ | |
| B, N, C, H, W = volume.shape | |
| x = volume.view(B * N, C, H, W) # flatten slices | |
| # (B*N, 3, 224, 224) → (B*N, embed_dim) → reshape | |
| slice_feats = self._slice_embed(x).view(B, N, -1) # (B, 32, embed_dim) | |
| # ---------------- ABMIL attention ---------------- # | |
| A = torch.tanh(self.attn_V(slice_feats)) # (B, 32, hidden) | |
| A = self.attn_U(A) # (B, 32, 1) | |
| A = torch.softmax(A, dim=1) # attention weights | |
| # Weighted sum → patient embedding | |
| patient_feat = (A * slice_feats).sum(dim=1) # (B, embed_dim) | |
| # ---------------- Head ---------------- # | |
| out = self.classifier(self.dropout(self.norm(patient_feat))) # (B, num_classes) | |
| return out, A.squeeze(-1) # logits, attention weights | |
| def main(args): | |
| read_excel = pd.read_csv(args.configuration_file, skip_blank_lines=True, na_values=['NaN']) | |
| max_f1_scores = [] | |
| model_files = [] | |
| for index, row in read_excel.iterrows(): | |
| if args.log_wandb: | |
| wandb.login() | |
| run = wandb.init(project=args.project_name, reinit=True, entity="adarshbhandary") | |
| with run: | |
| print('Experiment Index', index) | |
| torch.cuda.empty_cache() | |
| print("Reading Configuration File..") | |
| config = { | |
| "cv_run": row['no'], | |
| "task": row['task'], | |
| "model": row['model'], | |
| "method": row['method'], | |
| "view": row['view'], | |
| "pretrained": row['pretrained'], | |
| "loss": row['loss'], | |
| "optimizer": row['optimizer'], | |
| "use_sampler": row['use_sampler'], | |
| "height": int(row['height']), | |
| "width": int(row['width']), | |
| "background_crop": row['background_crop'], | |
| "use_clahe": row['use_clahe'], | |
| "batch_size": int(row['batch_size']), | |
| "multi_gpu": bool(row['multi_gpu']), | |
| "num_worker": row['num_worker'], | |
| "learning_rate": row['lr'], | |
| "weight_decay": row['weight_decay'], | |
| "drop_out": row['drop_out'], | |
| "num_epochs": int(row['epochs']), | |
| "patience": row['patience'], | |
| "probability": row['probability'], | |
| "aug_mix_p":row['aug_mix_p'], | |
| "erasing":row['erasing'], | |
| "attention_head": row["attention_head"], | |
| "rage":row["rage"], | |
| "get_gradcam": bool(row['get_gradcam']), | |
| "out_folder": row['out_folder'], | |
| "folder_name": row['folder_name'], | |
| "excel_file_train": row['train'], | |
| "excel_file_validation": row['valid'], | |
| "excel_file_test": row['test'], | |
| "data_folder": row['data_folder'], | |
| "pretrained_weights": row["pretrained_weights"], | |
| "results_path": row['csv_results'] | |
| } | |
| set_seed(int(config["cv_run"])) | |
| wandb.config.update(config, allow_val_change=True) | |
| os.makedirs(os.path.join(config["out_folder"], str(config["folder_name"])), exist_ok=True) | |
| save_out_folder = config["out_folder"] + str(config["folder_name"]) | |
| ## config method is MRI | |
| target_names = extract_target_names(method = config["method"]) | |
| num_classes = len(target_names) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print('Device : ', device) | |
| weight, samples_weight = compute_sample_weights(config["excel_file_train"], class_type = config["method"], view = config["view"]) | |
| if config["use_sampler"]: | |
| sampler = WeightedRandomSampler(samples_weight, num_samples=len(samples_weight), replacement=True) | |
| shuffle = False | |
| else: | |
| sampler = None | |
| shuffle = True | |
| if config["model"] == "vit_small": | |
| train_transforms = TrainTransform(aug=True, aug_mix_p = config["aug_mix_p"]) | |
| else: | |
| train_transforms = TrainTransformBaseline(aug=True, aug_mix_p = config["aug_mix_p"], erasing = config["erasing"]) | |
| valid_transforms = BasicTransforms() | |
| if config["method"] == "MRI_MIL": | |
| data_train = BreastMRI_ABMIL_NII( | |
| split='train' | |
| ) | |
| data_valid = BreastMRI_ABMIL_NII( | |
| split='val' | |
| ) | |
| data_test = BreastMRI_ABMIL_NII( | |
| split='test' | |
| ) | |
| else: | |
| data_train = BreastMRI(csv_file = config["excel_file_train"], | |
| root_dir=None, | |
| split='train', | |
| transform = valid_transforms) | |
| data_valid = BreastMRI(csv_file = config["excel_file_train"], | |
| root_dir=None, | |
| split='val', | |
| transform = valid_transforms) | |
| data_test = BreastMRI(csv_file = config["excel_file_train"], | |
| root_dir=None, | |
| split='test', | |
| transform = valid_transforms) | |
| print("Train dataset size:", len(data_train)) | |
| print("Val dataset size:", len(data_valid)) | |
| print("Test dataset size:", len(data_test)) | |
| dataloader = { | |
| 'train': torch.utils.data.DataLoader(data_train, | |
| batch_size=config["batch_size"], | |
| shuffle=shuffle, | |
| sampler=sampler, | |
| num_workers=config["num_worker"], | |
| pin_memory=True, | |
| prefetch_factor=2, | |
| drop_last=True | |
| ), | |
| 'valid': torch.utils.data.DataLoader(data_valid, | |
| batch_size=1, | |
| shuffle=False, | |
| sampler=None, | |
| num_workers=config["num_worker"], | |
| pin_memory=True, | |
| prefetch_factor=2, | |
| drop_last=False | |
| ), | |
| 'test': torch.utils.data.DataLoader(data_test, | |
| batch_size=1, | |
| shuffle=False, | |
| sampler=None, | |
| num_workers=config["num_worker"], | |
| pin_memory=True, | |
| prefetch_factor=2, | |
| drop_last=False | |
| ) | |
| } | |
| if config["method"]=="MRI": | |
| if config["model"] == "swin_cross": | |
| model = MultiViewSwinCrossAttn(num_classes=num_classes) | |
| else: | |
| model = MultiViewSwinModel(num_classes=num_classes) | |
| elif config["method"] == "MRI_MIL": | |
| if config["model"] == "swin_cross": | |
| model = CrossModalAttentionABMIL_Swin(num_classes=num_classes) | |
| else: | |
| model = ABMIL_Swin(num_classes=num_classes) | |
| model = model.to(device) | |
| #summary(model, (config["batch_size"], 3, config["height"], config["width"])) | |
| if config["loss"] == "F": | |
| criterion = FocalLoss(gamma=config["erasing"], reduction='mean').to(device) | |
| elif config["loss"] == "CE": | |
| criterion = nn.CrossEntropyLoss().to(device) | |
| optimizer_ft = optim.Adam(model.parameters(), | |
| lr=config["learning_rate"], | |
| weight_decay=config["weight_decay"]) | |
| if config["method"] == 'MRI': | |
| train = FullySupervisedMultiViewTrainer(device, | |
| config["method"], | |
| model, | |
| criterion, | |
| optimizer_ft, | |
| dataloader, | |
| target_names, | |
| save_out_folder, | |
| config["num_epochs"], | |
| config["patience"]) | |
| train.main_loop() | |
| model.load_state_dict(torch.load(save_out_folder + 'max_metrics_epoch.pth')) | |
| train = FullySupervisedMultiViewTrainer(device, | |
| config["method"], | |
| model, | |
| criterion, | |
| optimizer_ft, | |
| dataloader, | |
| target_names, | |
| save_out_folder, | |
| config["num_epochs"], | |
| config["patience"]) | |
| test_accuracy, test_f1, true, pred, test_auc, averaged_results = train.test_loop() | |
| elif config["method"] == 'MRI_MIL': | |
| #model.load_state_dict(torch.load(save_out_folder + 'max_metrics_epoch.pth')) | |
| train = MultiInstanceTrainer(device, | |
| config["method"], | |
| model, | |
| criterion, | |
| config["loss"], | |
| optimizer_ft, | |
| dataloader, | |
| target_names, | |
| save_out_folder, | |
| config["num_epochs"], | |
| config["patience"]) | |
| train.main_loop() | |
| model.load_state_dict(torch.load(save_out_folder + 'max_metrics_epoch.pth')) | |
| train = MultiInstanceTrainer(device, | |
| config["method"], | |
| model, | |
| criterion, | |
| config["loss"], | |
| optimizer_ft, | |
| dataloader, | |
| target_names, | |
| save_out_folder, | |
| config["num_epochs"], | |
| config["patience"]) | |
| test_accuracy, test_f1, true, pred, test_auc, averaged_results = train.test_loop() | |
| print('Test Results') | |
| print(classification_report(true, pred, target_names=target_names, digits=4)) | |
| display_confusion_matrix(true, | |
| pred, | |
| target_names, | |
| save_out_folder+'confusion_matrix.png', | |
| use_tta=False, | |
| label=config["task"], | |
| normalize=False) | |
| config["model location"] = save_out_folder + 'least_validation_loss.pth' | |
| config["test accuracy"] = test_accuracy | |
| config["test f1"] = test_f1 | |
| config["test auc"] = test_auc | |
| #log_scalar_values(config["results_path"], config) | |
| max_f1_scores.append(test_f1) | |
| model_files.append(save_out_folder + 'least_validation_loss.pth') | |
| wandb.log({ | |
| "test_accuracy" : test_accuracy, | |
| "test_f1" : test_f1, | |
| "test_auc" : test_auc, | |
| "test_metrics" : averaged_results, | |
| }) | |
| #log_scalar_values(config["results_path"], config) | |
| wandb.config.update(config, allow_val_change=True) | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--configuration_file", type=Path, default='/cluster/panambur/DINO/config/transfer_baseline.csv', | |
| help="""Change hyperparameters and input the filepath of CSV file""") | |
| parser.add_argument("--log_wandb", type=bool, default=True, | |
| help="""Use FALSE for Siemens Code.""") | |
| parser.add_argument("--project_name", type=str, default='BreastMRI Tiny', | |
| help="""For Wandb logging""") | |
| args = parser.parse_args() | |
| main(args) |