File size: 1,891 Bytes
e8ae657
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
import argparse
import torch
from src.training.config import TrainingConfig
from src.training.trainer import Trainer


def parse_args():
    parser = argparse.ArgumentParser(
        description="Train the multi-task CV pipeline"
    )
    parser.add_argument(
        "--epochs", type=int, default=30,
        help="Number of training epochs"
    )
    parser.add_argument(
        "--batch-size", type=int, default=32,
        help="Batch size"
    )
    parser.add_argument(
        "--lr", type=float, default=1e-3,
        help="Learning rate"
    )
    parser.add_argument(
        "--lambda-cls", type=float, default=1.0,
        help="Classification loss weight"
    )
    parser.add_argument(
        "--lambda-det", type=float, default=5.0,
        help="Detection loss weight"
    )
    parser.add_argument(
        "--max-train-samples", type=int, default=None,
        help="Cap training samples (None = full dataset)"
    )
    parser.add_argument(
        "--max-val-samples", type=int, default=None,
        help="Cap validation samples (None = full dataset)"
    )
    parser.add_argument(
        "--experiment-name", type=str, default="multitask_v1",
        help="Name for this training run"
    )
    parser.add_argument(
        "--voc-root", type=str, default="data/VOCdevkit/VOC2012",
        help="Path to VOC2012 root"
    )
    return parser.parse_args()


def main():
    args = parse_args()

    config = TrainingConfig(
        voc_root=args.voc_root,
        num_epochs=args.epochs,
        batch_size=args.batch_size,
        learning_rate=args.lr,
        lambda_cls=args.lambda_cls,
        lambda_det=args.lambda_det,
        max_train_samples=args.max_train_samples,
        max_val_samples=args.max_val_samples,
        experiment_name=args.experiment_name,
    )

    trainer = Trainer(config)
    trainer.train()


if __name__ == "__main__":
    main()