eel4759_classification

Comparison of different ImageNet-based CNN models for classifying diabetic retinopathy images
Log | Files | Refs | README

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:
Mdataset.py | 6++++--
Mmain.py | 4+++-
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)