3d_model / docs /ADVANCED_OPTIMIZATIONS.md
Azan
Clean deployment build (Squashed)
7a87926
|
Raw History Blame Contribute Delete
19.6 kB

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
  2. Advanced Training Techniques
  3. Inference Optimizations
  4. Data Pipeline Enhancements
  5. System-Level Optimizations
  6. Memory Optimizations

Model Compilation & Optimization

1. Torch Compile (PyTorch 2.0+)

Impact: 1.5-3x faster training/inference, minimal code changes

Implementation:

# 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:

# 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:

# 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:

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:

# 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:

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:

# 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:

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:

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:

# 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:

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:

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:

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:

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:

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:

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:

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:

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:

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:

# 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:

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:

# 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)

  1. βœ… Batch inference
  2. βœ… Async data loading
  3. βœ… HDF5 datasets
  4. βœ… Gradient checkpointing (if needed)

Phase 3: Advanced (1-2 weeks)

  1. βœ… DDP for multi-GPU
  2. βœ… Model quantization
  3. βœ… ONNX/TensorRT export
  4. βœ… 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:

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.