CoolFace
Apppublic

Droid210/FleetVision

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
main.py127 linesDownload Raw Back to model_b
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