commit 3d5c9e1658ea8c8b06d9acaeff847647b58fa63a
parent 9dd8b3007e79f0634f868d405f5e7c48917305fd
Author: Vineet Kumar <git@vineetk.net>
Date: Mon, 20 Apr 2026 10:28:06 -0400
increase resolution to 384px and soften oversampler weights
- Default img_size changed from 224 to 384 (already plumbed through
build_transforms/build_dataloaders, just needed passing from main.py)
- WeightedRandomSampler now uses sqrt(1/count) instead of 1/count to
reduce over-prediction of minority classes against visually similar
majority class (Healthy vs Mild NPDR)
- Default batch size reduced 64->32 to accommodate larger images
- 4-epoch warmup-only test: macro F1 0.45, weighted F1 0.74, kappa 0.62
already exceeding previous 22-epoch best on all metrics
Diffstat:
2 files changed, 7 insertions(+), 3 deletions(-)
diff --git a/dataset.py b/dataset.py
@@ -126,10 +126,12 @@ def build_dataloaders(
val_ds = DRDataset(val_paths, val_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
test_ds = DRDataset(test_paths, test_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
- # Per-sample weights: inverse of class frequency so all classes are seen equally
+ # Per-sample weights: sqrt-inverse-frequency for moderate class balancing.
+ # Full inverse (1/count) causes over-prediction of minorities; sqrt gives a
+ # softer boost that retains enough majority-class signal.
labels_arr = np.array(train_labels)
class_counts = np.bincount(labels_arr, minlength=5)
- class_sample_weights = 1.0 / class_counts.astype(float)
+ class_sample_weights = 1.0 / np.sqrt(class_counts.astype(float))
sample_weights = class_sample_weights[labels_arr]
sampler = WeightedRandomSampler(
weights=torch.from_numpy(sample_weights).float(),
diff --git a/main.py b/main.py
@@ -37,7 +37,7 @@ def parse_args():
)
parser.add_argument('--epochs', type=int, default=50, help='Max training epochs')
parser.add_argument('--patience', type=int, default=10, help='Early stopping patience')
- parser.add_argument('--batch-size', type=int, default=64, help='Training batch size')
+ parser.add_argument('--batch-size', type=int, default=32, help='Training batch size')
parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate')
parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay')
parser.add_argument('--seed', type=int, default=42, help='Random seed')
@@ -46,6 +46,7 @@ def parse_args():
parser.add_argument('--num-workers', type=int, default=8, help='DataLoader workers')
parser.add_argument('--focal-gamma', type=float, default=2.0, help='Focal loss gamma (0=standard CE)')
parser.add_argument('--warmup-epochs', type=int, default=3, help='LR warmup epochs')
+ parser.add_argument('--img-size', type=int, default=384, help='Input image size')
return parser.parse_args()
@@ -92,6 +93,7 @@ def main():
use_clahe=use_clahe,
batch_size=args.batch_size,
num_workers=args.num_workers,
+ img_size=args.img_size,
)
model = create_model(model_name, num_classes=5, pretrained=True).to(device)