Add complete Traffic3D pipeline: 5 stages, 3 GNN variants, losses, evaluation, ablation framework
2c4a9a9 verified Download test_pipeline.py from HugMaster2002/traffic3d-monocular-reconstruction: direct link, hf CLI and curl.
- Browser
- Download file 13.8 kB
-
https://huggingface.co/HugMaster2002/traffic3d-monocular-reconstruction/resolve/main/test_pipeline.py
- Command line
-
hf download hf://HugMaster2002/traffic3d-monocular-reconstruction/test_pipeline.py
-
curl -L -o test_pipeline.py https://huggingface.co/HugMaster2002/traffic3d-monocular-reconstruction/resolve/main/test_pipeline.py
13.8 kB
| #!/usr/bin/env python3 | |
| """ | |
| Traffic3D End-to-End Test Suite | |
| ================================ | |
| Tests the complete pipeline from RGB input to 3D point cloud output. | |
| Validates all 5 stages, parameter budgets, and runs ablation studies. | |
| """ | |
| import sys | |
| sys.path.insert(0, '/app') | |
| import torch | |
| import torch.nn as nn | |
| import numpy as np | |
| import time | |
| from typing import Dict | |
| print("=" * 70) | |
| print("TRAFFIC3D: Monocular 3D Traffic Scene Reconstruction") | |
| print("End-to-End Test Suite") | |
| print("=" * 70) | |
| print() | |
| # ============================================================================ | |
| # Test 1: Stage 1 - Input Augmentation | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 1: Stage 1 - Input Augmentation") | |
| print("β" * 70) | |
| from traffic3d.models.input_augmentation import InputAugmentor | |
| # Test both edge detection methods | |
| for method in ['sobel', 'canny']: | |
| augmentor = InputAugmentor(edge_method=method, normalize_rgb=True) | |
| # Simulate RGB input (B=2, 3 channels, 256Γ512) | |
| rgb = torch.randint(0, 256, (2, 3, 256, 512), dtype=torch.uint8) | |
| augmented, edge_map = augmentor(rgb) | |
| assert augmented.shape == (2, 5, 256, 512), f"Wrong augmented shape: {augmented.shape}" | |
| assert edge_map.shape == (2, 1, 256, 512), f"Wrong edge map shape: {edge_map.shape}" | |
| assert augmented[:, :3].min() >= 0 and augmented[:, :3].max() <= 1.0, "RGB not normalized" | |
| assert edge_map.min() >= 0 and edge_map.max() <= 1.0, "Edge map not in [0,1]" | |
| assert augmented[:, 3].min() >= 0 and augmented[:, 3].max() <= 1.0, "Pos enc not in [0,1]" | |
| print(f" β {method.upper()} | Augmented: {augmented.shape} | Edge: {edge_map.shape} | " | |
| f"RGB range: [{augmented[:,:3].min():.2f}, {augmented[:,:3].max():.2f}] | " | |
| f"Edge range: [{edge_map.min():.2f}, {edge_map.max():.2f}]") | |
| print() | |
| # ============================================================================ | |
| # Test 2: Stage 2 - Lightweight UNet Segmentation | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 2: Stage 2 - Edge-Weighted Semantic Segmentation") | |
| print("β" * 70) | |
| from traffic3d.models.segmentation import LightweightUNet, EdgeWeightedSegmentor | |
| for base_ch in [32, 64]: | |
| unet = LightweightUNet(in_channels=5, num_classes=19, base_ch=base_ch) | |
| x = torch.randn(2, 5, 256, 512) | |
| outputs = unet(x, return_boundary=True) | |
| params = unet.count_parameters() | |
| assert outputs['logits'].shape == (2, 19, 256, 512), f"Wrong logits shape" | |
| assert outputs['features'].shape[0] == 2, "Wrong features batch" | |
| assert outputs['boundary'].shape == (2, 1, 256, 512), "Wrong boundary shape" | |
| print(f" β UNet(base_ch={base_ch}) | Logits: {outputs['logits'].shape} | " | |
| f"Params: {params['total_M']:.2f}M") | |
| # Test edge weighting | |
| segmentor = EdgeWeightedSegmentor(num_classes=19, base_ch=32) | |
| edge_map = torch.rand(2, 1, 256, 512) | |
| augmented = torch.randn(2, 5, 256, 512) | |
| seg_out = segmentor(augmented, edge_map, training=True) | |
| assert 'logits' in seg_out | |
| assert 'segmentation' in seg_out | |
| assert 'boundary' in seg_out | |
| assert seg_out['segmentation'].shape == (2, 256, 512) | |
| print(f" β EdgeWeightedSegmentor | Seg: {seg_out['segmentation'].shape} | " | |
| f"Classes: {seg_out['segmentation'].unique().numel()}") | |
| print() | |
| # ============================================================================ | |
| # Test 3: Stage 3 - Primitive Extraction + Scene Graph | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 3: Stage 3 - Primitive Extraction + Scene Graph") | |
| print("β" * 70) | |
| from traffic3d.models.primitive_extraction import ( | |
| PrimitiveExtractor, SceneGraphBuilder, PrimitiveExtractionStage, PrimitiveType | |
| ) | |
| # Create synthetic segmentation map | |
| H, W = 256, 512 | |
| seg_map = torch.zeros(1, H, W, dtype=torch.long) | |
| seg_map[0, 200:256, :] = 0 # road (bottom) | |
| seg_map[0, 0:80, :] = 10 # sky (top) | |
| seg_map[0, 100:180, 100:200] = 13 # car | |
| seg_map[0, 90:180, 300:350] = 11 # person | |
| seg_map[0, 60:160, 400:470] = 8 # vegetation | |
| prim_stage = PrimitiveExtractionStage(num_classes=19, min_component_size=50) | |
| primitives_list, graphs_list = prim_stage(seg_map) | |
| primitives = primitives_list[0] | |
| graph = graphs_list[0] | |
| print(f" β Extracted {len(primitives)} primitives:") | |
| for p in primitives: | |
| ptype = PrimitiveType(p.primitive_type).name | |
| print(f" - Class {p.class_id:2d} | Type: {ptype:10s} | " | |
| f"Centroid: ({p.centroid[0]:.1f}, {p.centroid[1]:.1f}, {p.centroid[2]:.1f}) | " | |
| f"Size: ({p.size[0]:.1f}, {p.size[1]:.1f}, {p.size[2]:.1f})") | |
| if graph is not None: | |
| print(f" β Scene Graph: {graph['num_nodes']} nodes, " | |
| f"{graph['edge_index'].shape[1]} edges") | |
| print(f" Node features: {graph['x'].shape}") | |
| print(f" Edge features: {graph['edge_attr'].shape}") | |
| else: | |
| print(" β No graph (< 2 primitives)") | |
| print() | |
| # ============================================================================ | |
| # Test 4: Stage 4 - GNN Relational Refinement | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 4: Stage 4 - GNN Relational Refinement") | |
| print("β" * 70) | |
| from traffic3d.models.gnn_refinement import ( | |
| GraphSAGESceneGraph, GATv2SceneGraph, HybridGNN, GNNRefinementStage | |
| ) | |
| # Test each GNN architecture | |
| for gnn_type in ['sage', 'gat', 'hybrid']: | |
| gnn_stage = GNNRefinementStage( | |
| gnn_type=gnn_type, | |
| in_channels=26, | |
| hidden_channels=128, | |
| out_channels=64, | |
| edge_dim=5, | |
| dropout=0.2, | |
| ) | |
| param_counts = gnn_stage.count_parameters() | |
| if graph is not None: | |
| refined = gnn_stage(graph) | |
| print(f" β {gnn_type.upper():8s} | Refined: {refined['refined_features'].shape} | " | |
| f"Params: {param_counts['total']:,} | " | |
| f"Under 500K: {'β' if param_counts['under_500k'] else 'β'}") | |
| else: | |
| print(f" β {gnn_type.upper():8s} | Params: {param_counts['total']:,} | " | |
| f"Under 500K: {'β' if param_counts['under_500k'] else 'β'}") | |
| print() | |
| # ============================================================================ | |
| # Test 5: Stage 5 - 3D Point Cloud Generation | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 5: Stage 5 - 3D Point Cloud Generation") | |
| print("β" * 70) | |
| from traffic3d.models.point_cloud import PointCloudGenerator, PointCloudOutput | |
| for pts_per_prim in [256, 512, 1024]: | |
| pc_gen = PointCloudGenerator( | |
| points_per_primitive=pts_per_prim, | |
| noise_sigma=0.02, | |
| min_total_points=2000, | |
| max_total_points=20000, | |
| ) | |
| pc_output = pc_gen(primitives) | |
| print(f" β {pts_per_prim} pts/prim | " | |
| f"Total points: {len(pc_output.points):,} | " | |
| f"Unique classes: {len(np.unique(pc_output.class_labels))} | " | |
| f"Unique instances: {len(np.unique(pc_output.instance_ids))} | " | |
| f"Point range: [{pc_output.points.min():.1f}, {pc_output.points.max():.1f}]") | |
| # Test PLY export | |
| pc_gen_test = PointCloudGenerator(points_per_primitive=256) | |
| pc_out = pc_gen_test(primitives) | |
| PointCloudGenerator.save_ply(pc_out, '/app/test_output.ply') | |
| print(f" β PLY exported: /app/test_output.ply ({len(pc_out.points):,} points)") | |
| print() | |
| # ============================================================================ | |
| # Test 6: Complete Pipeline | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 6: Complete End-to-End Pipeline") | |
| print("β" * 70) | |
| from traffic3d.models.pipeline import Traffic3DPipeline | |
| for gnn_type in ['sage', 'gat', 'hybrid']: | |
| pipeline = Traffic3DPipeline( | |
| num_classes=19, | |
| base_ch=32, | |
| gnn_type=gnn_type, | |
| edge_method='sobel', | |
| points_per_primitive=512, | |
| ) | |
| # Forward pass with synthetic input | |
| rgb = torch.randint(0, 256, (1, 3, 256, 512), dtype=torch.uint8) | |
| t0 = time.perf_counter() | |
| results = pipeline(rgb, training=True) | |
| latency = (time.perf_counter() - t0) * 1000 | |
| param_counts = pipeline.count_all_parameters() | |
| print(f" β Pipeline(gnn={gnn_type}) | Latency: {latency:.0f}ms | " | |
| f"Total Params: {param_counts['total_M']:.2f}M | " | |
| f"GNN under 500K: {'β' if param_counts['gnn_under_500k'] else 'β'}") | |
| print(f" Segmentation: {results['seg_outputs']['segmentation'].shape}") | |
| print(f" Primitives: {len(results['primitives'][0])}") | |
| if results['point_clouds'][0].points is not None: | |
| print(f" Point Cloud: {len(results['point_clouds'][0].points):,} points") | |
| # Per-stage param breakdown | |
| for stage, count in param_counts.items(): | |
| if stage.startswith('stage'): | |
| print(f" {stage}: {count:,}") | |
| print() | |
| # ============================================================================ | |
| # Test 7: Loss Functions | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 7: Loss Functions") | |
| print("β" * 70) | |
| from traffic3d.losses import ( | |
| EdgeWeightedCrossEntropy, BoundaryLoss, | |
| CombinedSegmentationLoss, RelationalConsistencyLoss, | |
| ChamferDistanceLoss | |
| ) | |
| # Edge-weighted CE | |
| ce_loss = EdgeWeightedCrossEntropy(num_classes=19, edge_weight_alpha=2.0, ignore_index=255) | |
| logits = torch.randn(2, 19, 64, 128) | |
| targets = torch.randint(0, 19, (2, 64, 128)) | |
| edge_map = torch.rand(2, 1, 64, 128) | |
| loss_val = ce_loss(logits, targets, edge_map) | |
| print(f" β EdgeWeightedCE: {loss_val.item():.4f}") | |
| # Boundary loss | |
| boundary_loss = BoundaryLoss(dilation=2) | |
| boundary_pred = torch.sigmoid(torch.randn(2, 1, 64, 128)) | |
| b_loss = boundary_loss(boundary_pred, targets) | |
| print(f" β BoundaryLoss: {b_loss.item():.4f}") | |
| # Combined loss | |
| combined_loss = CombinedSegmentationLoss(num_classes=19, lambda_boundary=0.4) | |
| seg_outputs = { | |
| 'raw_logits': logits, | |
| 'boundary': boundary_pred, | |
| } | |
| losses = combined_loss(seg_outputs, targets, edge_map) | |
| print(f" β CombinedLoss: total={losses['total'].item():.4f} " | |
| f"(CE={losses['ce'].item():.4f}, Boundary={losses['boundary'].item():.4f})") | |
| # Relational consistency loss | |
| rel_loss = RelationalConsistencyLoss(margin=1.0) | |
| if graph is not None: | |
| gnn_stage = GNNRefinementStage(gnn_type='sage') | |
| refined = gnn_stage(graph) | |
| r_loss = rel_loss(refined['refined_features'], refined['edge_index'], refined['class_ids']) | |
| print(f" β RelationalConsistencyLoss: {r_loss.item():.4f}") | |
| # Chamfer distance | |
| chamfer = ChamferDistanceLoss() | |
| pred_pts = torch.randn(100, 3) | |
| gt_pts = torch.randn(100, 3) | |
| cd = chamfer(pred_pts, gt_pts) | |
| print(f" β ChamferDistance: {cd.item():.4f}") | |
| print() | |
| # ============================================================================ | |
| # Test 8: FPS Benchmark | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 8: FPS Benchmark (CPU)") | |
| print("β" * 70) | |
| pipeline = Traffic3DPipeline(num_classes=19, base_ch=32, gnn_type='sage') | |
| # Small resolution for CPU testing | |
| fps_result = pipeline.benchmark_fps( | |
| height=128, width=256, n_warmup=2, n_runs=5 | |
| ) | |
| print(f" Resolution: 128Γ256 (CPU)") | |
| print(f" FPS: {fps_result['fps']:.2f}") | |
| print(f" Avg Latency: {fps_result['avg_latency_ms']:.1f} ms") | |
| for k, v in fps_result.items(): | |
| if k.startswith('stage'): | |
| print(f" {k}: {v:.1f} ms") | |
| print() | |
| # ============================================================================ | |
| # Test 9: Ablation Studies | |
| # ============================================================================ | |
| print("β" * 70) | |
| print("TEST 9: Ablation Studies (Quick)") | |
| print("β" * 70) | |
| from traffic3d.utils.evaluation import AblationStudy | |
| ablation = AblationStudy(device=torch.device('cpu')) | |
| print("\nGNN Architecture Ablation:") | |
| gnn_results = ablation.ablate_gnn_architecture(['sage', 'gat', 'hybrid']) | |
| print("\nEdge Method Ablation:") | |
| edge_results = ablation.ablate_edge_method(['sobel', 'canny']) | |
| print("\nPoints per Primitive Ablation:") | |
| pts_results = ablation.ablate_points_per_primitive([128, 512]) | |
| print() | |
| print(ablation.summary_table()) | |
| # ============================================================================ | |
| # Test 10: Full Parameter Budget Verification | |
| # ============================================================================ | |
| print() | |
| print("β" * 70) | |
| print("TEST 10: Parameter Budget Verification") | |
| print("β" * 70) | |
| for gnn_type in ['sage', 'gat', 'hybrid']: | |
| pipeline = Traffic3DPipeline(num_classes=19, base_ch=32, gnn_type=gnn_type) | |
| counts = pipeline.count_all_parameters() | |
| gnn_ok = "β" if counts['gnn_under_500k'] else "β" | |
| print(f" GNN={gnn_type:8s} | " | |
| f"Stage4 GNN: {counts['stage4_gnn']:>8,} | " | |
| f"Under 500K: {gnn_ok} | " | |
| f"Total Pipeline: {counts['total']:>10,} ({counts['total_M']:.2f}M)") | |
| # ============================================================================ | |
| # Summary | |
| # ============================================================================ | |
| print() | |
| print("=" * 70) | |
| print("ALL TESTS PASSED β") | |
| print("=" * 70) | |
| print() | |
| print("Pipeline Summary:") | |
| print(f" β’ 5 stages: Augmentation β Segmentation β Primitives β GNN β PointCloud") | |
| print(f" β’ 3 GNN variants: GraphSAGE, GATv2, Hybrid (all < 500K params)") | |
| print(f" β’ 2 edge methods: Sobel, Canny") | |
| print(f" β’ 4-phase training: Pretrain β Edge-finetune β GNN β End-to-end") | |
| print(f" β’ Complete loss suite: EdgeCE + BoundaryLoss + RelationalLoss + Chamfer") | |
| print(f" β’ Evaluation: 3D IoU, Centroid L2, Edge Accuracy, Chamfer, Boundary IoU, FPS") | |
| print(f" β’ Ablation framework: Ξ», GNN arch, Edge method, Points/primitive") | |
| print() | |