# Advanced Training & Inference Optimizations This document outlines advanced optimization techniques beyond the basic improvements, targeting 5-10x additional speedups and better training stability. ## Table of Contents 1. [Model Compilation & Optimization](#model-compilation--optimization) 2. [Advanced Training Techniques](#advanced-training-techniques) 3. [Inference Optimizations](#inference-optimizations) 4. [Data Pipeline Enhancements](#data-pipeline-enhancements) 5. [System-Level Optimizations](#system-level-optimizations) 6. [Memory Optimizations](#memory-optimizations) --- ## Model Compilation & Optimization ### 1. Torch Compile (PyTorch 2.0+) **Impact**: 1.5-3x faster training/inference, minimal code changes **Implementation**: ```python # In model_loader.py def load_da3_model(..., compile_model: bool = True): model = DepthAnything3.from_pretrained(model_name) model = model.to(device) if compile_model and hasattr(torch, 'compile'): logger.info("Compiling model with torch.compile...") # Compile for inference model = torch.compile(model, mode="reduce-overhead", fullgraph=False) # For training, use mode="max-autotune" or "default" model.eval() return model # In training loops, compile forward pass if use_compile: model_forward = torch.compile(model.forward, mode="reduce-overhead") else: model_forward = model.forward ``` **Benefits**: - Automatic kernel fusion - Better GPU utilization - Works with existing code **Caveats**: - First run is slower (compilation overhead) - Some dynamic operations may not compile ### 2. cuDNN Benchmark Mode **Impact**: 10-30% faster convolutions **Implementation**: ```python # At start of training script if torch.backends.cudnn.is_available(): torch.backends.cudnn.benchmark = True # Optimize for consistent input sizes torch.backends.cudnn.deterministic = False # Allow non-deterministic for speed ``` **When to use**: - Input sizes are consistent - Training (not inference where determinism matters) ### 3. JIT Compilation for Custom Operations **Impact**: 2-5x faster custom loss functions **Implementation**: ```python # In losses.py @torch.jit.script def geodesic_rotation_loss_jit(R_pred: torch.Tensor, R_target: torch.Tensor) -> torch.Tensor: R_diff = torch.matmul(R_pred, R_target.transpose(-2, -1)) trace = torch.diagonal(R_diff, dim1=-2, dim2=-1).sum(dim=-1) trace_clamped = torch.clamp(trace, -1.0, 3.0) angle = torch.acos((trace_clamped - 1.0) / 2.0) return angle.mean() ``` --- ## Advanced Training Techniques ### 4. Exponential Moving Average (EMA) **Impact**: Better model stability, improved final performance **Implementation**: ```python class EMA: def __init__(self, model, decay=0.9999): self.model = model self.decay = decay self.shadow = {} self.backup = {} self.register() def register(self): for name, param in self.model.named_parameters(): if param.requires_grad: self.shadow[name] = param.data.clone() def update(self): for name, param in self.model.named_parameters(): if param.requires_grad: assert name in self.shadow new_average = (1.0 - self.decay) * param.data + self.decay * self.shadow[name] self.shadow[name] = new_average.clone() def apply_shadow(self): for name, param in self.model.named_parameters(): if param.requires_grad: assert name in self.shadow self.backup[name] = param.data param.data = self.shadow[name] def restore(self): for name, param in self.model.named_parameters(): if param.requires_grad: assert name in self.backup param.data = self.backup[name] self.backup = {} # In training loop ema = EMA(model, decay=0.9999) for batch in dataloader: # ... training step ... ema.update() # Update EMA after each step # Use EMA model for evaluation ema.apply_shadow() eval_loss = evaluate(model, val_loader) ema.restore() ``` **Benefits**: - Smoother training dynamics - Better generalization - More stable checkpoints ### 5. Gradient Checkpointing **Impact**: 40-60% memory reduction, 20-30% slower (trade-off) **Implementation**: ```python # For models that support it from torch.utils.checkpoint import checkpoint class CheckpointedModel(nn.Module): def forward(self, x): # Checkpoint intermediate layers x = checkpoint(self.layer1, x) x = checkpoint(self.layer2, x) return x # Or use activation checkpointing in training if use_gradient_checkpointing: model.gradient_checkpointing_enable() ``` **When to use**: - Running out of memory - Large models - Can trade speed for memory ### 6. Learning Rate Finder / OneCycleLR **Impact**: Faster convergence, better final performance **Implementation**: ```python from torch.optim.lr_scheduler import OneCycleLR # Replace CosineAnnealingLR with OneCycleLR scheduler = OneCycleLR( optimizer, max_lr=lr * 10, # Peak LR (10x base) epochs=epochs, steps_per_epoch=len(dataloader), pct_start=0.1, # 10% warmup anneal_strategy='cos', div_factor=10.0, # Initial LR = max_lr / div_factor final_div_factor=100.0, # Final LR = max_lr / final_div_factor ) ``` **Benefits**: - Automatically finds good learning rate - Superconvergence training - Better than manual LR scheduling ### 7. Label Smoothing **Impact**: Better generalization, reduced overfitting **Implementation**: ```python # In loss computation def smooth_pose_loss(poses_pred, poses_target, smoothing=0.1): # Add small noise to targets noise = torch.randn_like(poses_target) * smoothing poses_target_smooth = poses_target + noise return pose_loss(poses_pred, poses_target_smooth) ``` ### 8. Focal Loss for Hard Examples **Impact**: Better focus on difficult samples **Implementation**: ```python def focal_pose_loss(poses_pred, poses_target, alpha=0.25, gamma=2.0): base_loss = pose_loss(poses_pred, poses_target) # Focus more on hard examples focal_weight = (base_loss / base_loss.max()) ** gamma return alpha * focal_weight * base_loss ``` --- ## Inference Optimizations ### 9. Batch Inference **Impact**: 2-5x faster when processing multiple sequences **Current Problem**: Model inference is called per-sequence **Implementation**: ```python class BatchedInference: def __init__(self, model, batch_size=4): self.model = model self.batch_size = batch_size self.queue = [] def add(self, images, sequence_id): self.queue.append((images, sequence_id)) if len(self.queue) >= self.batch_size: return self.process_batch() return None def process_batch(self): # Batch all images together all_images = [] sequence_boundaries = [] idx = 0 for images, seq_id in self.queue: all_images.extend(images) sequence_boundaries.append((idx, idx + len(images))) idx += len(images) # Run batched inference with torch.no_grad(): outputs = self.model.inference(all_images) # Split results back results = [] for (start, end), (_, seq_id) in zip(sequence_boundaries, self.queue): result = { 'extrinsics': outputs.extrinsics[start:end], 'intrinsics': outputs.intrinsics[start:end] if hasattr(outputs, 'intrinsics') else None, 'sequence_id': seq_id, } results.append(result) self.queue = [] return results ``` ### 10. Model Quantization (INT8/FP16) **Impact**: 2-4x faster inference, 50-75% memory reduction **Implementation**: ```python # Post-training quantization def quantize_model(model, calibration_data): model.eval() model_fp16 = model.half() # FP16 quantization # Or INT8 quantization (more complex) model_int8 = torch.quantization.quantize_dynamic( model, {torch.nn.Linear, torch.nn.Conv2d}, dtype=torch.qint8 ) return model_int8 # Use quantized model for inference quantized_model = quantize_model(model, calibration_loader) ``` **When to use**: - Inference-only workloads - Memory-constrained environments - Production deployments ### 11. ONNX/TensorRT Export **Impact**: 3-10x faster inference on optimized runtimes **Implementation**: ```python def export_to_onnx(model, sample_input, output_path): model.eval() torch.onnx.export( model, sample_input, output_path, input_names=['images'], output_names=['extrinsics', 'intrinsics', 'depth'], dynamic_axes={ 'images': {0: 'batch_size'}, 'extrinsics': {0: 'batch_size'}, }, opset_version=17, ) # Then use ONNX Runtime or TensorRT for inference import onnxruntime as ort session = ort.InferenceSession("model.onnx") outputs = session.run(None, {"images": input_numpy}) ``` ### 12. Inference Caching **Impact**: Instant results for repeated queries **Implementation**: ```python from functools import lru_cache import hashlib class CachedInference: def __init__(self, model, cache_dir=None): self.model = model self.cache = {} self.cache_dir = cache_dir def _hash_images(self, images): # Create hash from image content combined = np.concatenate([img.flatten()[:1000] for img in images]) return hashlib.md5(combined.tobytes()).hexdigest() def inference(self, images): cache_key = self._hash_images(images) if cache_key in self.cache: return self.cache[cache_key] result = self.model.inference(images) self.cache[cache_key] = result return result ``` --- ## Data Pipeline Enhancements ### 13. Async Data Loading **Impact**: Eliminate data loading bottlenecks **Implementation**: ```python from torch.utils.data import DataLoader import asyncio from concurrent.futures import ThreadPoolExecutor class AsyncDataLoader: def __init__(self, dataloader, prefetch=2): self.dataloader = dataloader self.prefetch = prefetch self.executor = ThreadPoolExecutor(max_workers=prefetch) self.queue = asyncio.Queue(maxsize=prefetch) async def _prefetch_worker(self): for batch in self.dataloader: await self.queue.put(batch) await self.queue.put(None) # Sentinel async def __aiter__(self): task = asyncio.create_task(self._prefetch_worker()) while True: batch = await self.queue.get() if batch is None: break yield batch await task ``` ### 14. Memory-Mapped Files (HDF5) **Impact**: Faster I/O, lower memory usage for large datasets **Implementation**: ```python import h5py class HDF5Dataset(Dataset): def __init__(self, hdf5_path): self.hdf5_path = hdf5_path self.file = h5py.File(hdf5_path, 'r') self.length = len(self.file['images']) def __getitem__(self, idx): # Memory-mapped access (no full load) images = self.file['images'][idx] poses = self.file['poses'][idx] return {'images': images, 'poses': poses} def __len__(self): return self.length # Create HDF5 file from existing data def create_hdf5_dataset(samples, output_path): with h5py.File(output_path, 'w') as f: images_ds = f.create_dataset('images', shape=(len(samples), N, H, W, 3), dtype=np.uint8) poses_ds = f.create_dataset('poses', shape=(len(samples), N, 3, 4), dtype=np.float32) for i, sample in enumerate(samples): images_ds[i] = np.stack(sample['images']) poses_ds[i] = sample['poses'] ``` ### 15. Smart Sampling (Curriculum Learning) **Impact**: Faster convergence, better final performance **Implementation**: ```python class CurriculumSampler: def __init__(self, dataset, difficulty_fn): self.dataset = dataset self.difficulty_fn = difficulty_fn # Function that scores sample difficulty self.weights = self._compute_weights() def _compute_weights(self): # Start with easy samples, gradually include harder ones difficulties = [self.difficulty_fn(sample) for sample in self.dataset] # Weight by inverse difficulty early, then uniform weights = 1.0 / (np.array(difficulties) + 1e-6) return weights def sample(self, epoch, total_epochs): # Gradually shift from easy to hard progress = epoch / total_epochs current_weights = self.weights * (1 - progress) + np.ones_like(self.weights) * progress return np.random.choice(len(self.dataset), p=current_weights/current_weights.sum()) ``` ### 16. Advanced Augmentation **Impact**: Better generalization, data efficiency **Implementation**: ```python import albumentations as A # Strong augmentation pipeline augmentation = A.Compose([ A.RandomBrightnessContrast(p=0.5), A.RandomGamma(p=0.3), A.GaussNoise(p=0.2), A.MotionBlur(p=0.2), A.OpticalDistortion(p=0.2), A.GridDistortion(p=0.2), # Geometric augmentations (be careful with poses!) # A.HorizontalFlip(p=0.5), # Only if poses are adjusted ]) # MixUp augmentation def mixup_data(x, y, alpha=1.0): lam = np.random.beta(alpha, alpha) index = torch.randperm(x.size(0)) mixed_x = lam * x + (1 - lam) * x[index] y_a, y_b = y, y[index] return mixed_x, y_a, y_b, lam # CutMix def cutmix_data(x, y, alpha=1.0): lam = np.random.beta(alpha, alpha) index = torch.randperm(x.size(0)) bbx1, bby1, bbx2, bby2 = rand_bbox(x.size(), lam) x[:, :, bbx1:bbx2, bby1:bby2] = x[index, :, bbx1:bbx2, bby1:bby2] y_a, y_b = y, y[index] return x, y_a, y_b, lam ``` --- ## System-Level Optimizations ### 17. Distributed Data Parallel (DDP) **Impact**: Linear scaling with number of GPUs **Implementation**: ```python import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP def setup_ddp(rank, world_size): dist.init_process_group("nccl", rank=rank, world_size=world_size) torch.cuda.set_device(rank) def train_ddp(rank, world_size, ...): setup_ddp(rank, world_size) model = load_da3_model(...) model = DDP(model, device_ids=[rank]) # Each process gets subset of data sampler = torch.utils.data.distributed.DistributedSampler( dataset, num_replicas=world_size, rank=rank ) dataloader = DataLoader(dataset, sampler=sampler, ...) # Training loop (same as before) for epoch in range(epochs): sampler.set_epoch(epoch) # Shuffle differently each epoch for batch in dataloader: # ... training ... ``` ### 18. Fully Sharded Data Parallel (FSDP) **Impact**: Train models that don't fit on single GPU **Implementation**: ```python from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import ShardingStrategy model = FSDP( model, sharding_strategy=ShardingStrategy.FULL_SHARD, mixed_precision=MixedPrecision( param_dtype=torch.float16, reduce_dtype=torch.float16, ), ) ``` ### 19. GPU/CPU Pipeline Parallelism **Impact**: Better utilization, hide CPU bottlenecks **Current Problem**: GPU waits for CPU (BA validation) **Implementation**: ```python from queue import Queue from threading import Thread class PipelineProcessor: def __init__(self, model, ba_validator, gpu_queue, cpu_queue): self.model = model self.ba_validator = ba_validator self.gpu_queue = gpu_queue self.cpu_queue = cpu_queue def gpu_worker(self): while True: item = self.gpu_queue.get() if item is None: break images, seq_id = item with torch.no_grad(): output = self.model.inference(images) self.cpu_queue.put((output, images, seq_id)) def cpu_worker(self): while True: item = self.cpu_queue.get() if item is None: break output, images, seq_id = item result = self.ba_validator.validate(images, output.extrinsics) # Process result... # Run GPU and CPU work in parallel gpu_thread = Thread(target=processor.gpu_worker) cpu_thread = Thread(target=processor.cpu_worker) gpu_thread.start() cpu_thread.start() ``` --- ## Memory Optimizations ### 20. Gradient Accumulation with Async **Impact**: Better GPU utilization during accumulation **Current**: Synchronous accumulation **Implementation**: ```python # Use async operations during accumulation async def async_backward(loss): loss.backward() # Do other work while backward is running await asyncio.sleep(0) # Yield to other tasks ``` ### 21. Dynamic Batch Sizing **Impact**: Maximize GPU utilization, avoid OOM **Implementation**: ```python class DynamicBatchSampler: def __init__(self, dataset, initial_batch_size=1, max_batch_size=8): self.dataset = dataset self.batch_size = initial_batch_size self.max_batch_size = max_batch_size self.oom_count = 0 def __iter__(self): try: # Try current batch size yield self._get_batch() except RuntimeError as e: if "out of memory" in str(e): # Reduce batch size on OOM self.batch_size = max(1, self.batch_size // 2) torch.cuda.empty_cache() yield self._get_batch() else: raise def on_success(self): # Gradually increase batch size if successful if self.oom_count == 0: self.batch_size = min(self.max_batch_size, self.batch_size * 2) self.oom_count = 0 ``` ### 22. Activation Offloading **Impact**: Trade compute for memory **Implementation**: ```python # Offload activations to CPU during forward pass class ActivationOffload(nn.Module): def forward(self, x): # Store on CPU, move to GPU when needed x = x.cpu() # ... compute ... x = x.cuda() return x ``` --- ## Implementation Priority ### Phase 1: Quick Wins (1-2 days) 1. ✅ Torch compile 2. ✅ cuDNN benchmark mode 3. ✅ EMA 4. ✅ OneCycleLR ### Phase 2: High Impact (3-5 days) 5. ✅ Batch inference 6. ✅ Async data loading 7. ✅ HDF5 datasets 8. ✅ Gradient checkpointing (if needed) ### Phase 3: Advanced (1-2 weeks) 9. ✅ DDP for multi-GPU 10. ✅ Model quantization 11. ✅ ONNX/TensorRT export 12. ✅ Pipeline parallelism --- ## Expected Combined Performance With all optimizations: - **Training speed**: 5-15x faster (depending on hardware) - **Inference speed**: 10-50x faster (with quantization/TensorRT) - **Memory usage**: 50-80% reduction - **GPU utilization**: 95-99% - **Scalability**: Linear with number of GPUs --- ## Monitoring & Profiling Add profiling to identify bottlenecks: ```python from torch.profiler import profile, record_function, ProfilerActivity with profile( activities=[ProfilerActivity.CUDA, ProfilerActivity.CPU], record_shapes=True, profile_memory=True, ) as prof: with record_function("training_step"): # Training code... print(prof.key_averages().table(sort_by="cuda_time_total")) ``` Use this to identify which optimizations will have the most impact for your specific workload.