Droid210/FleetVision
0
1"""CLI entry point for damage detector training."""2import argparse3import os4from pathlib import Path5 6from .config import TrainConfig7from .data import build_dataloaders8from .evaluate import evaluate_model9from .model import build_model10from .train import train_model11from .utils import get_device, set_seed12 13 14def parse_args() -> TrainConfig:15 """Parse command-line arguments.16 17 Returns:18 TrainConfig object.19 """20 root_dir = Path(__file__).resolve().parents[2]21 22 parser = argparse.ArgumentParser(description="Train ViT damage detector.")23 parser.add_argument(24 "--data-dir",25 type=Path,26 default=root_dir / "data" / "model b" / "damage_assessment",27 help=(28 "Dataset root. Supported layouts: "29 "(1) train/valid folders with Whole/Damaged subfolders, or "30 "(2) damage_assessment format (data/ + samples.json) with Whole images "31 "available in whole_pool/, Whole/, train/Whole/, or valid/Whole/."32 ),33 )34 parser.add_argument("--batch-size", type=int, default=32)35 parser.add_argument("--epochs", type=int, default=15)36 parser.add_argument("--learning-rate", type=float, default=1e-4)37 parser.add_argument("--num-workers", type=int, default=2)38 parser.add_argument(39 "--model-output",40 type=Path,41 default=root_dir / "weights" / "model b" / "best_damage_detector.pth",42 )43 parser.add_argument("--seed", type=int, default=42)44 parser.add_argument("--recall-weight", type=float, default=2.0)45 46 args = parser.parse_args()47 return TrainConfig(48 data_dir=args.data_dir,49 batch_size=args.batch_size,50 epochs=args.epochs,51 learning_rate=args.learning_rate,52 num_workers=args.num_workers,53 model_output=args.model_output,54 seed=args.seed,55 recall_weight=args.recall_weight,56 )57 58 59def main() -> None:60 """Main training pipeline."""61 config = parse_args()62 63 # On Windows, multiprocessing workers can fail under mixed Python envs.64 if os.name == "nt" and config.num_workers > 0:65 print("WARNING: Forcing num_workers=0 on Windows for stable DataLoader execution.")66 config.num_workers = 067 68 set_seed(config.seed)69 70 if not config.data_dir.exists():71 raise FileNotFoundError(f"Data directory not found: {config.data_dir}")72 73 device = get_device()74 print(f"Using device: {device}")75 76 print("\n=== Loading Data ===")77 train_loader, val_loader, processor = build_dataloaders(78 data_dir=config.data_dir,79 batch_size=config.batch_size,80 num_workers=config.num_workers,81 )82 print(f"Training samples: {len(train_loader.dataset)}")83 print(f"Validation samples: {len(val_loader.dataset)}")84 85 print("\n=== Building Model ===")86 model = build_model().to(device)87 88 print("\n=== Training (Recall-Focused) ===")89 train_model(90 model=model,91 train_loader=train_loader,92 val_loader=val_loader,93 epochs=config.epochs,94 learning_rate=config.learning_rate,95 device=device,96 output_path=config.model_output,97 recall_weight=config.recall_weight,98 )99 100 print("\n=== Validation Set Evaluation ===")101 evaluate_model(102 model=model,103 loader=val_loader,104 device=device,105 )106 107 print("\n=== Inference Example ===")108 print(109 "from models.model_b.inference import classify_damage\n"110 "class_name, confidence = classify_damage('path/to/car_image.jpg', "111 f"model_path='{config.model_output}')\n"112 "print(f'Status: {class_name}, Confidence: {confidence:.2%}')"113 )114 115 print("\n=== Multi-View Inspection Example ===")116 print(117 "from models.model_b.inspection import inspect_vehicle_sync\n"118 "images = ['front.jpg', 'back.jpg', 'left.jpg', 'right.jpg', 'roof.jpg']\n"119 "verdict, results = inspect_vehicle_sync(images, "120 f"model_path='{config.model_output}')\n"121 "print(f'Verdict: {verdict}')"122 )123 124 125if __name__ == "__main__":126 main()127 