summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--dataset.py6
-rw-r--r--main.py4
2 files changed, 7 insertions, 3 deletions
diff --git a/dataset.py b/dataset.py
index 0068a97..a5a955b 100644
--- a/dataset.py
+++ b/dataset.py
@@ -126,10 +126,12 @@ def build_dataloaders(
126 val_ds = DRDataset(val_paths, val_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size)) 126 val_ds = DRDataset(val_paths, val_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
127 test_ds = DRDataset(test_paths, test_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size)) 127 test_ds = DRDataset(test_paths, test_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
128 128
129 # Per-sample weights: inverse of class frequency so all classes are seen equally 129 # Per-sample weights: sqrt-inverse-frequency for moderate class balancing.
130 # Full inverse (1/count) causes over-prediction of minorities; sqrt gives a
131 # softer boost that retains enough majority-class signal.
130 labels_arr = np.array(train_labels) 132 labels_arr = np.array(train_labels)
131 class_counts = np.bincount(labels_arr, minlength=5) 133 class_counts = np.bincount(labels_arr, minlength=5)
132 class_sample_weights = 1.0 / class_counts.astype(float) 134 class_sample_weights = 1.0 / np.sqrt(class_counts.astype(float))
133 sample_weights = class_sample_weights[labels_arr] 135 sample_weights = class_sample_weights[labels_arr]
134 sampler = WeightedRandomSampler( 136 sampler = WeightedRandomSampler(
135 weights=torch.from_numpy(sample_weights).float(), 137 weights=torch.from_numpy(sample_weights).float(),
diff --git a/main.py b/main.py
index fb87c33..6c4cd20 100644
--- a/main.py
+++ b/main.py
@@ -37,7 +37,7 @@ def parse_args():
37 ) 37 )
38 parser.add_argument('--epochs', type=int, default=50, help='Max training epochs') 38 parser.add_argument('--epochs', type=int, default=50, help='Max training epochs')
39 parser.add_argument('--patience', type=int, default=10, help='Early stopping patience') 39 parser.add_argument('--patience', type=int, default=10, help='Early stopping patience')
40 parser.add_argument('--batch-size', type=int, default=64, help='Training batch size') 40 parser.add_argument('--batch-size', type=int, default=32, help='Training batch size')
41 parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate') 41 parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate')
42 parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay') 42 parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay')
43 parser.add_argument('--seed', type=int, default=42, help='Random seed') 43 parser.add_argument('--seed', type=int, default=42, help='Random seed')
@@ -46,6 +46,7 @@ def parse_args():
46 parser.add_argument('--num-workers', type=int, default=8, help='DataLoader workers') 46 parser.add_argument('--num-workers', type=int, default=8, help='DataLoader workers')
47 parser.add_argument('--focal-gamma', type=float, default=2.0, help='Focal loss gamma (0=standard CE)') 47 parser.add_argument('--focal-gamma', type=float, default=2.0, help='Focal loss gamma (0=standard CE)')
48 parser.add_argument('--warmup-epochs', type=int, default=3, help='LR warmup epochs') 48 parser.add_argument('--warmup-epochs', type=int, default=3, help='LR warmup epochs')
49 parser.add_argument('--img-size', type=int, default=384, help='Input image size')
49 return parser.parse_args() 50 return parser.parse_args()
50 51
51 52
@@ -92,6 +93,7 @@ def main():
92 use_clahe=use_clahe, 93 use_clahe=use_clahe,
93 batch_size=args.batch_size, 94 batch_size=args.batch_size,
94 num_workers=args.num_workers, 95 num_workers=args.num_workers,
96 img_size=args.img_size,
95 ) 97 )
96 98
97 model = create_model(model_name, num_classes=5, pretrained=True).to(device) 99 model = create_model(model_name, num_classes=5, pretrained=True).to(device)