# Tested shape for the GPU port. Pin these; transformers in particular is load-bearing # because the script relies on cu_seq_lens_q/k reaching flash_attn_varlen_func # (transformers.modeling_flash_attention_utils, "Case 2" padding-free path). --extra-index-url https://download.pytorch.org/whl/cu128 torch==2.8.0 transformers==4.57.5 tokenizers>=0.22 datasets>=3.1.0,<4.0.0 # v4 dropped dataset-script support; also what the marin eval branch pins accelerate>=1.0.0 safetensors>=0.4.5 numpy>=1.26 jinja2>=3.1.0 # only needed by prepare_sft_data.py --verify sentencepiece huggingface_hub>=0.26 # FlashAttention-2 is REQUIRED for the 32K varlen path. Build takes ~20 min, or use the # prebuilt wheel matching your torch/cuda/python/abi from the flash-attention releases page. flash-attn==2.8.3 --no-build-isolation # optional wandb>=0.18