diff --git a/DGX_SPARK_SETUP.md b/DGX_SPARK_SETUP.md index 8c1a9ac..c7511f0 100644 --- a/DGX_SPARK_SETUP.md +++ b/DGX_SPARK_SETUP.md @@ -155,24 +155,47 @@ training, so all-zero/all-text is correct): - the reference-model forward pass, `self.ref_model is None` branch (inside `null_ref_context()`) - the reference-model forward pass, `self.ref_model is not None` branch -## 7. FTPO fine-tuning is compute-bound, not throughput-bound +## 7. FTPO fine-tuning was slow because of fixed-length padding, not raw compute With `finetune_batch_size: 1` / `gradient_accumulation_steps: 16`, we measured ~750 optimizer steps at ~160-185s/step — a ~34 hour run for the full 12,000-example dataset. Bumping `finetune_batch_size` to 4 (with `gradient_accumulation_steps` dropped to 4 to keep the same effective batch size) made **no meaningful difference** — still ~160-175s/step. -Takeaway: this workload is compute-bound (two full forward passes per micro-batch — the model plus -a reference-model pass for the MSE tether loss term — over sequences up to -`finetune_max_seq_length: 4000` tokens), not limited by batch-size/scheduling overhead. Increasing -batch size doesn't reduce total FLOPs for a fixed effective batch size, so it doesn't help here. -If you need a faster run, the actual levers are: -- lower `finetune_max_train_examples` (fewer total steps, less data coverage) -- lower `finetune_max_seq_length` (less compute per step, truncates longer training examples) -- accept the long runtime and let it run in the background +Initial hypothesis was that this was inherent — genuinely compute-bound (two full forward passes +per micro-batch, over long sequences), not limited by batch-size/scheduling overhead. Measuring the +actual training data disproved that. `ftpo_trainer.py`'s collator pads every batch to a **fixed** +`finetune_max_seq_length` (4000 tokens) regardless of content: -We left `finetune_batch_size` at the default (`1`) since increasing it only costs more memory for -no speed benefit on this hardware. +```python +max_len = self.args.max_length # always 4000, never pad-to-longest-in-batch +prompt_ids = torch.full((batch_sz, max_len), pad_id, dtype=torch.long) +``` + +We tokenized all 12,000 training contexts with the real tokenizer to see how much of that 4000 was +actually needed: + +| | tokens | +|---|---| +| mean | 529.9 | +| median | 509 | +| p90 / p99 | 953 / 1080 | +| **max across all 12,000 examples** | **1126** | + +Not one example reaches even a third of the 4000-token padding target; the mean uses 13.2% of it. +Every forward pass — both the main model and the reference-model pass — was processing ~4000 tokens +of mostly padding, roughly 4-7x more than the actual content needs. This also explains why the +batch-size bump did nothing: total padded-token compute is invariant to how the effective batch of +16 gets split into micro-batches, so reshuffling batch/accum never touched the real cost. This is a +collator-design issue, not a hardware ceiling — it would waste the same proportion on any GPU. + +**Fix:** lower `finetune_max_seq_length` to comfortably cover the real distribution, e.g. `1280` +(covers p99 with headroom, nothing in the dataset gets truncated) instead of `4000`. That should cut +per-step compute roughly 3x, bringing the ~34h estimate down to somewhere around ~11-12h. We left +`finetune_batch_size` at the default (`1`) since increasing it has no effect either way here. + +If you need it faster still, the other lever is `finetune_max_train_examples` (fewer total steps, +less data coverage) — or just accept the runtime and let it run in the background. ## Validated results diff --git a/configs/gemma-3-4b-it.yaml b/configs/gemma-3-4b-it.yaml index ec280f9..e3d9b54 100644 --- a/configs/gemma-3-4b-it.yaml +++ b/configs/gemma-3-4b-it.yaml @@ -210,7 +210,7 @@ finetune_mode: "ftpo" # ftpo / dpo / dpo_final_token finetune_ftpo_dataset: "" # you can specify an existing ftpo dataset, or leave unset to let the # pipeline use the one produced in the generation step finetune_base_model_id: null # Base model for DPO (if unset, uses model_id) -finetune_max_seq_length: 4000 # this may truncate some outputs +finetune_max_seq_length: 1280 # measured p99 context length is 1080 tokens (max observed: 1126) -- 4000 was mostly wasted padding (~13% utilization), ~3x more compute per step than needed on this dataset finetune_load_in_4bit: true # qlora # --- Early Stopping ---