Model Training Infrastructure and Distributed Training Questions
Scaling model training across hardware and time. Covers GPU/accelerator considerations, data and model parallelism, distributed and large-scale training, experiment tracking and training infrastructure, and the training-versus-inference compute tradeoff. Focuses on the systems and resource decisions that make large-model training feasible.
You need to implement a custom fused transformer-attention kernel to leverage Tensor Cores for better throughput. Explain design choices and implementation steps either using CUDA WMMA APIs or Triton: data layout, tile sizes, alignment constraints, memory staging (shared memory), avoiding bank conflicts, and validation for numerical correctness and performance regression testing.
Sample Answer
Direct answer
Implementing a custom fused transformer-attention kernel to leverage Tensor Cores means combining several steps of the attention computation (the QK^T matrix multiply, the softmax, and the subsequent multiply by V) into a single GPU kernel that keeps intermediate results in fast on-chip memory rather than round-tripping through HBM between each step, while structuring the matrix multiplications specifically to match Tensor Cores' preferred tile sizes and precision.
Structured elaboration
- Why fusion matters here specifically: a naive, unfused attention implementation computes QK^T (writing the full attention score matrix to HBM), then reads it back for softmax (writing the result back to HBM again), then reads it back again for the final multiply by V; for long sequences, this attention-score matrix is large (quadratic in sequence length) and this round-tripping is a substantial, avoidable HBM bandwidth cost. Fusing these steps into one kernel (as FlashAttention-style implementations do) keeps intermediate results in fast shared memory/registers, computing the whole attention operation for a tile of the sequence without ever materializing the full attention-score matrix in HBM.
- Tiling for Tensor Cores: the QK^T and attention-weights-times-V matrix multiplications need to be structured as tiled operations matching Tensor Cores' preferred small-matrix-tile granularity (Tensor Cores operate on fixed small tile sizes, e.g. 16x16 or similar, depending on generation and precision), requiring the kernel to explicitly manage how the sequence and head dimensions are partitioned into Tensor-Core-sized tiles, rather than treating the operation as one large, unstructured matrix multiply.
- Precision choices: using fp16/bf16 for the matrix multiplications (to engage Tensor Cores) while keeping the softmax's numerically sensitive operations (the exponentiation and normalization) in fp32 internally, converting back to the lower precision only for the final output, balancing Tensor Core throughput against the numerical stability softmax specifically needs.
- Online softmax (a key algorithmic trick): since the full attention-score matrix is never materialized, softmax needs to be computed incrementally (an "online" or "streaming" softmax) as tiles of the score matrix are produced, maintaining a running maximum and running sum that get corrected as new, potentially-larger values are seen, rather than the standard softmax algorithm's assumption that the full row of values is available at once.
- Alignment constraints: tile dimensions and shared-memory buffer strides need to be sized and aligned to the hardware's requirements (e.g. global-memory accesses aligned to 128-byte boundaries for efficient coalescing, and matrix dimensions padded to the Tensor Core's tile-size multiple, as discussed in the Tensor Core tile-shapes topic) so that both the global-memory loads staging data into shared memory and the Tensor Core operations themselves run on their fastest supported path rather than an unaligned, slower fallback.
- Avoiding shared-memory bank conflicts: shared memory is physically divided into banks, and when multiple threads in the same warp read or write different addresses that happen to map to the same bank simultaneously, those accesses serialize instead of completing in parallel, a "bank conflict" that silently degrades the kernel's effective shared-memory bandwidth; staging Q/K/V tiles into shared memory needs a layout that avoids this, typically via padding (adding an extra unused column so a tile's row stride no longer aligns with the bank count) or an explicit swizzled/permuted addressing scheme, a technique used deliberately in FlashAttention-style kernels to keep every thread in a warp hitting a distinct bank.
Worked example
Design choices in code (framework-level, since a full custom CUDA kernel is out of scope for this format): implement attention using PyTorch's torch.nn.functional.scaled_dot_product_attention with a Flash-Attention-style backend selected (which internally performs exactly this fusion and Tensor-Core-tiled computation), or, if building a fully custom kernel, use a library like Triton to write the fused kernel at a higher level of abstraction than raw CUDA while still controlling tile sizes and the online-softmax algorithm explicitly; either path avoids materializing the full O(sequence_length^2) attention-score matrix in HBM, which is the core design win regardless of implementation level.
Trade-offs & pitfalls
Writing a genuinely optimal custom fused attention kernel from scratch (rather than using an existing, heavily-optimized implementation like FlashAttention) is a substantial engineering undertaking, requiring careful tuning of tile sizes for the specific hardware generation and getting the online-softmax numerics exactly right; for most teams, the practical answer is using an existing, well-validated fused-attention implementation rather than reimplementing this from scratch, reserving custom kernel work for cases with genuinely novel attention variants not covered by existing libraries. Validating a custom kernel also needs its own explicit process: check numerical correctness by comparing the kernel's output against a reference eager-mode (unfused) attention implementation using a tolerance-based comparison (e.g. torch.allclose with an atol/rtol appropriate to the reduced precision in use, since fp16/bf16 accumulation will not match fp32 bit-for-bit) across a range of sequence lengths, batch sizes, and edge cases (very short sequences, causal/masked attention, sequences not evenly divisible by the tile size); and set up a performance regression test that benchmarks the kernel's measured throughput/latency against a stored baseline on every change, failing the check if performance regresses beyond a defined threshold, since a kernel change can stay numerically correct while silently becoming slower.
When scaling a training job from 8 to 64 GPUs you observe only 2x speedup rather than ~8x. Provide a rigorous debugging plan to identify the bottleneck including specific metrics to collect, profiling tools to use for compute and network, experiments to isolate I/O versus compute versus communication, and potential remediation steps.
Sample Answer
Direct answer
Observing only 2x speedup instead of the expected roughly 8x when scaling from 8 to 64 GPUs points to communication overhead (or another form of non-compute overhead) increasingly dominating as GPU count grows; a rigorous debugging plan isolates whether the bottleneck is communication, data loading, or a configuration issue that silently changed between the two scales, rather than assuming poor "scaling efficiency" is an unavoidable fact of distributed training.
Structured elaboration
- Step 1, profile at both scales and compare: capture a profiler trace (per-GPU compute time versus communication time) at both 8-GPU and 64-GPU configurations; if communication's share of step time grows substantially from 8 to 64 GPUs (a common and expected pattern, since all-reduce communication volume per worker grows, though sub-linearly, with worker count), that directly quantifies how much of the shortfall is communication-attributable. Use concrete tools for this: a framework profiler (
torch.profileror NVIDIA Nsight Systems) for the per-GPU compute/kernel timeline,nvidia-smi/dcgmfor real-time utilization and memory telemetry,NCCL_DEBUG=INFOto confirm the actual transport and topology NCCL selected, andnccl-tests(e.g.all_reduce_perf) run in isolation to measure the interconnect's achieved all-reduce bandwidth independent of the training job, giving a network-only baseline to compare the in-job communication time against. - Step 1b, isolate I/O (data loading) as a candidate cause separately from compute and communication: run the same data loader alone (no model forward/backward) at both 8 and 64 GPUs and measure samples/sec it can sustain; if the data pipeline's aggregate throughput doesn't scale with worker count (e.g. because it reads from a shared filesystem or a single-process data source that becomes a bottleneck as more GPU workers request data concurrently), GPUs at 64 scale may sit idle waiting on data even if compute and communication are both healthy, a cause distinct from anything the compute/communication profiler split alone would catch.
- Step 2, check per-GPU batch size and effective utilization: if global batch size was held constant while scaling from 8 to 64 GPUs, per-GPU batch size shrank 8x, which can itself hurt per-GPU compute efficiency (smaller batches are less efficient to compute, independent of any distributed-training-specific issue) even before considering communication overhead; check whether per-GPU batch size at 64 GPUs is still large enough to keep the GPU's compute units well-utilized.
- Step 3, check the interconnect topology being used at 64 GPUs: confirm the 64-GPU configuration actually spans multiple nodes with a network topology that supports it well (correct NCCL configuration, RDMA enabled if available, no accidental fallback to a much slower communication path); a misconfiguration here (e.g. NCCL silently falling back to a slower transport) can cause exactly this kind of scaling shortfall and is a common, checkable root cause.
- Step 4, check for a straggler or heterogeneous-hardware issue: confirm all 64 GPUs are genuinely homogeneous and none is a straggler (per the earlier straggler-detection discussion), since synchronous training's step time is bounded by the slowest participant, and a single consistently slow node among 64 can disproportionately drag down overall speedup.
- Step 5, compute and report scaling efficiency explicitly: 2x speedup for 8x more GPUs is 25% scaling efficiency; report this number explicitly (not just "it's slower than expected") and use it as the metric to track as each hypothesis above is tested and, hopefully, fixed.
Worked example
Profiling reveals communication's share of step time grew from roughly 10% at 8 GPUs to roughly 60% at 64 GPUs, immediately explaining the bulk of the shortfall: at 8 GPUs, compute dominates (efficient scaling), while at 64 GPUs, communication has become the dominant cost, consistent with the well-known pattern of all-reduce overhead growing (even if sub-linearly) with worker count on a fixed interconnect; combined with a check confirming per-GPU batch size at 64 GPUs (having shrunk 8x from the 8-GPU configuration) is still reasonably large, and confirming no straggler is present, the investigation concludes communication overhead at the larger scale is the dominant root cause, pointing toward the communication-optimization techniques (overlap, bucketing, potentially compression) discussed elsewhere in this topic as the next step, rather than a distributed-training bug.
Trade-offs & pitfalls
A common mistake is attributing poor scaling directly to "distributed training overhead" as an unavoidable fact without actually profiling to confirm communication (rather than, say, a straggler, a configuration regression, or shrunken per-GPU batch efficiency) is the real cause; the rigorous version of this debugging plan insists on measuring and isolating the actual bottleneck before reaching for a communication-specific fix that might not even be the right one.
Outline a CI pipeline for ML training code that runs unit tests, environment reproducibility checks, small-scale integration trainings, and artifact validation before allowing full-scale runs. Describe tools, gating criteria, and how you would prevent flaky non-deterministic behavior from failing CI.
Sample Answer
Direct answer
A CI pipeline for ML training code should run fast unit tests on every commit, environment reproducibility checks, small-scale integration training runs, and artifact validation before any change reaches full-scale training, with the whole pipeline deliberately designed around tolerant, seeded, threshold-based checks rather than exact-match assertions, since naive exact-equality checks on anything involving GPU-kernel or stochastic behavior will be flaky and erode the team's trust in CI.
Structured elaboration
- Unit tests: fast, isolated tests of individual components: a custom loss function's output on known inputs, a data-preprocessing function's correctness, a model architecture's output shape.
- Environment reproducibility checks: verify the dependency lockfile resolves to the exact expected versions.
- Small-scale integration trainings: run the actual training entry point for a handful of steps on a small synthetic or subsampled dataset, checking the loop executes without error, loss is finite and decreasing, and checkpoint save/load round-trips correctly.
- Artifact validation: validate that expected artifacts were produced with expected structure/schema.
- Preventing flaky non-deterministic behavior from failing CI: this needs to be an explicit design principle across all four tiers, not an afterthought. First, every CI training run is seeded (fixed RNG seeds across Python/NumPy/framework, as in the general reproducibility checklist) so at least the SAME CI run is reproducible if re-triggered; deterministic-kernel flags are enabled specifically in CI (accepting the performance cost, since CI runs are tiny) so GPU-kernel-level non-determinism doesn't add noise on top of RNG seeding. Second, CI assertions on numerical outputs use TOLERANCE-based comparisons (e.g.
assert abs(loss - expected) < 1e-4, or a directional check like "loss decreased over N steps" rather than "loss equals exactly X"), since even with seeding and deterministic mode, minor floating-point differences across different CI runner hardware generations are still possible and a bit-exact assertion would fail spuriously on a legitimate hardware change. Third, a genuinely flaky check (one observed to fail intermittently despite no real code change) is treated as a bug in the TEST, not silently retried into passing or ignored: it's investigated (usually an under-seeded random source, exactly the failure mode a good reproducibility checklist is meant to prevent) and fixed or, if it can't be immediately fixed, explicitly quarantined (marked known-flaky, excluded from the merge-blocking gate, and tracked as an open issue) rather than left in the blocking suite where it silently trains the team to re-run CI without investigating red builds.
Worked example
A pull request modifying the data-augmentation pipeline triggers: unit tests confirming correctly-shaped output; an environment check confirming the lockfile resolves cleanly; and a 20-step integration training run on a 100-example synthetic dataset with a fixed seed and deterministic kernels enabled, asserting loss is finite and its value after step 20 is within a small tolerance band of a previously-recorded reference value (not required to match exactly), all completing in under 5 minutes. When this check failed intermittently once despite no code change, the team traced it to the augmentation library's own internally-seeded RNG not being re-seeded from the run's base seed between CI invocations; fixing that seeding gap (not adding a retry-until-green step) resolved the flakiness.
Trade-offs & pitfalls
The integration-training tier needs to be genuinely small and fast to be practical as a CI gate. A common design mistake specific to flakiness is reaching for automatic retries ("just re-run failed CI jobs once") as the default fix; this masks real flakiness sources (like the under-seeded augmentation RNG above) instead of fixing them, and a genuinely broken change can pass on a lucky retry, which is a worse outcome than an honest, investigated red build.
Design an experiment-tracking backend that stores large artifacts (models, datasets) and provides efficient search, lineage, and access control. Discuss object store integration, metadata DB schema, indexing for metric queries, caching for hot artifacts, and retention policies to balance cost and usability.
Sample Answer
Direct answer
An experiment-tracking backend at scale needs to separate metadata storage (a queryable database for run records, hyperparameters, and metric time series) from large-artifact storage (object storage for model checkpoints and datasets, referenced by pointer from the metadata layer, not stored inline), plus a lineage graph connecting runs to the data/code/artifacts they consumed and produced, and access control enforced at both the metadata and artifact-storage layers.
Structured elaboration
- Metadata store: a database (relational or a purpose-built time-series-friendly store for the metrics table specifically) holding run records, hyperparameters, and metric time series, optimized for the query patterns teams actually need (filter/sort runs by hyperparameter or metric value, retrieve a specific run's full history quickly).
- Artifact storage, separated: large binary artifacts (model checkpoints, potentially many GB each, and referenced datasets) live in object storage, with the metadata store holding only a reference (storage path/URI plus a content hash for integrity verification), never the artifact bytes themselves; this separation is what lets the metadata store stay fast and queryable even as total artifact storage grows into the petabyte range.
- Efficient search: indexing hyperparameters and key metrics for fast filtering (e.g. "find all runs with validation accuracy above X and learning rate in range Y") requires either a database with good support for semi-structured/JSON querying, or promoting frequently-queried hyperparameters to indexed, structured columns rather than leaving everything in an opaque JSON blob.
- Lineage: a graph (or graph-like relational structure) connecting each run to the exact data version it consumed, the code commit it ran, and the artifacts it produced, enabling both forward queries ("what runs used this dataset version") and backward queries ("what data/code produced this specific model artifact"), essential for debugging, auditing, and reproducibility investigations.
- Caching for hot artifacts: a small fraction of artifacts (the most recently trained checkpoints, or ones referenced by active experiments/dashboards) get accessed repeatedly while the vast majority of historical artifacts are accessed rarely if ever; fronting the object-storage tier with a caching layer (a CDN-style edge cache, or a local SSD cache on the compute nodes that most frequently need to load checkpoints for continued training or evaluation) for this hot subset avoids repeatedly paying object-storage egress cost and latency for the same frequently-reloaded artifacts, while leaving the cold, rarely-accessed majority in cheaper, uncached object storage.
- Retention policies: not every artifact needs to be kept forever; a retention policy (e.g. keep every checkpoint for the most recent N days or the most recent M runs in full, then downsample to only the final checkpoint per run for anything older, then eventually archive to cheaper cold storage or delete entirely per a defined age/relevance threshold) balances the real cost of storing every artifact indefinitely against the real usability cost of deleting something a team later needs; retention decisions are usually tiered by artifact importance (a promoted production model's artifacts retained far longer than an ordinary exploratory run's) rather than a single blanket policy.
- Access control: enforce permissions at the metadata layer (who can view/query which runs, often scoped by team or project) and separately at the artifact-storage layer (who can actually download a specific checkpoint, since metadata visibility and artifact-download rights are legitimately different permission levels, e.g. a run's existence and metrics might be broadly visible for collaboration while the actual trained model weights are more tightly restricted).
Worked example
A 50PB-scale artifact-storage tier holding checkpoints for tens of thousands of runs, referenced by a metadata database holding only lightweight pointers (a few hundred bytes per artifact reference, not the artifact itself); a query for "all runs from the last month with validation accuracy above 0.85" executes entirely against the (comparatively small, fast) metadata store's indexed columns, never touching the 50PB artifact tier at all, and only once a specific run of interest is identified does a separate, authorized request fetch the actual checkpoint bytes from object storage.
Trade-offs & pitfalls
Storing large artifacts inline in the same database as run metadata (rather than separating them into object storage with pointer references) is the most common early-stage design mistake, one that works fine at small scale but becomes a severe performance and cost problem as artifact volume grows, since it forces every metadata query's underlying storage engine to also manage petabyte-scale binary blob storage it was never designed for.
Analyze the convergence implications of stale gradients under asynchronous training. Quantify how staleness (delay s steps) can bias updates and propose algorithmic mitigations such as bounded staleness, learning rate adjustment, and momentum correction. Discuss how you would empirically validate whether staleness is causing divergence.
Sample Answer
Direct answer
Under asynchronous training with staleness of s steps, a worker's gradient is computed with respect to parameters that are s updates out of date; this biases the effective update direction (it's a gradient of an earlier point on the loss surface, applied as if it were current), and the magnitude of this bias grows with both staleness s and the learning rate, which is why mitigations focus on either bounding staleness or shrinking the effective step size for staler updates.
Structured elaboration
- Formalizing the bias: if θt is the parameter state a worker pulled and θt+s is the actual current state by the time its gradient ∇f(θt) is applied, the update effectively uses ∇f(θt) as an approximation of ∇f(θt+s); the error between these two gradients is (to first order) proportional to how far θt and θt+s have diverged, which in turn scales with staleness s and the learning rate η (a larger η means parameters move further per step, so the same staleness s corresponds to a larger parameter-space divergence).
- Algorithmic mitigations: (1) bounded staleness, rejecting or deferring any update whose staleness exceeds a threshold, directly caps the worst-case bias; (2) staleness-aware learning-rate adjustment, down-weighting (multiplying by a factor decreasing in s) a stale update's effective contribution, since a highly stale gradient is a less trustworthy estimate of the true current gradient direction; (3) momentum correction, adjusting how a stale update interacts with the optimizer's momentum term to avoid compounding the bias across successive stale updates.
- Empirical validation approach: to determine whether observed instability is actually caused by staleness (versus some other cause), instrument the training system to log each applied update's staleness value alongside the loss trajectory; if periods of high measured staleness (e.g. a straggler slowing down, causing others' effective staleness to rise) correlate with periods of slower loss decrease or increased loss variance, that's direct evidence staleness is a contributing factor, distinguishable from correlation-only intuition.
Worked example
With a fixed learning rate η=0.01 and observed staleness ranging from s=1 (near-synchronous) to s=10 (a significantly lagging worker) across a training run: applying a staleness-aware down-weighting of 1+s1 to each update's effective learning rate means an s=10 update contributes at roughly 111≈9% of its nominal weight, substantially limiting how much a highly stale (and therefore less trustworthy) gradient can distort the current parameter trajectory relative to a fresh, s=1 update contributing at roughly 50% of nominal weight.
Trade-offs & pitfalls
Down-weighting stale updates too aggressively effectively discards a straggler's computational contribution almost entirely, which somewhat defeats the point of choosing asynchronous training for utilization in the first place; the down-weighting function's steepness is itself a tunable trade-off between staleness-bias mitigation and making full use of every worker's compute, and should be validated empirically against the specific training run's observed staleness distribution rather than fixed a priori.
Unlock Full Question Bank
Get access to all Model Training Infrastructure and Distributed Training interview questions and detailed answers.
Sign in to ContinueJoin thousands of developers preparing for their dream job.