diff options
| -rw-r--r-- | dataset.py | 6 | ||||
| -rw-r--r-- | main.py | 4 |
2 files changed, 7 insertions, 3 deletions
| @@ -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(), |
| @@ -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) |
