HugMaster2002's picture
Add complete Traffic3D pipeline: 5 stages, 3 GNN variants, losses, evaluation, ablation framework
2c4a9a9 verified
Raw History Blame Contribute Delete
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()