LongStraw ਨੇ ਐਲਾਨ ਕੀਤਾ ਹੈ ਕਿ ਇਸਦੀ branch-replay ਤਕਨੀਕ ਸਿਰਫ਼ ਅੱਠ H20 GPUs ਦੀ ਵਰਤੋਂ ਕਰਕੇ reinforcement-learning (RL) post-training ਲਈ 2.1 ਮਿਲੀਅਨ token positions ਨੂੰ ਪ੍ਰੋਸੈਸ ਕਰ ਸਕਦੀ ਹੈ, ਜਿਸ ਨਾਲ ਹਾਰਡਵੇਅਰ ਦੇ ਖਰਚੇ ਵਿੱਚ ਕਈ ਗੁਣਾ ਦੀ ਕਮੀ ਆਉਂਦੀ ਹੈ। ਇਹ ਦਾਅਵਾ ਮਹੱਤਵਪੂਰਨ ਹੈ ਕਿਉਂਕਿ long-context ਮਾਡਲਾਂ ਦੀ ਟ੍ਰੇਨਿੰਗ ਲਈ ਰਵਾਇਤੀ ਤੌਰ 'ਤੇ ਦਰਜਨਾਂ ਉੱਚ-ਦਰਜੇ ਦੇ GPUs ਦੀ ਲੋੜ ਹੁੰਦੀ ਹੈ, ਜੋ ਕਿ ਜ਼ਿਆਦਾਤਰ ਰਿਸਰਚ ਲੈਬਾਂ ਅਤੇ ਸਟਾਰਟਅੱਪਸ ਲਈ ਇੱਕ ਰੁਕਾਵਟ ਹੈ।

long-context RL ਮਹਿੰਗਾ ਕਿਉਂ ਹੈ

Large language models ਦੀ RL-ਅਧਾਰਤ fine-tuning ਵਿੱਚ ਆਮ ਤੌਰ 'ਤੇ rollouts ਚਲਾਏ ਜਾਂਦੇ ਹਨ ਜੋ ਇੱਕੋ ਪ੍ਰੋਂਪਟ (prompt) ਲਈ ਕਈ ਵਿਕਲਪਿਕ ਮੁਕੰਮਲਤਾਵਾਂ (completions) ਪੈਦਾ ਕਰਦੇ ਹਨ। ਹਰੇਕ rollout ਨੂੰ back-propagate ਕਰਨਾ ਪੈਂਦਾ ਹੈ, ਇਸ ਲਈ compute cost ਪ੍ਰੋਸੈਸ ਕੀਤੇ ਗਏ token positions ਦੀ ਕੁੱਲ ਸੰਖਿਆ ਦੇ ਅਨੁਸਾਰ ਵਧਦੀ ਜਾਂਦੀ ਹੈ। ਮੌਜੂਦਾ pipelines ਜੋ ਇੱਕ ਮਿਲੀਅਨ-ਟੋਕਨ context ਦਾ ਟੀਚਾ ਰੱਖਦੇ ਹਨ, ਉਹਨਾਂ ਨੂੰ ਇੱਕ ਉਚਿਤ ਸਮੇਂ ਦੇ ਅੰਦਰ ਖਤਮ ਕਰਨ ਲਈ ਅਕਸਰ 64 ਤੋਂ 128 GPUs ਦੀ ਲੋੜ ਹੁੰਦੀ ਹੈ। ਉਸ ਹਾਰਡਵੇਅਰ ਦੀ ਕੀਮਤ, ਨਾਲ ਹੀ ਇਸ ਦੁਆਰਾ ਲੋੜੀਂਦੀ ਬਿਜਲੀ ਅਤੇ ਕੂਲਿੰਗ, ਇਹ ਸੀਮਤ ਕਰਦੀ ਹੈ ਕਿ ਅਭਿਆਸਕਰਤਾ (practitioners) context length ਨੂੰ ਕਿੰਨਾ ਵਧਾ ਸਕਦੇ ਹਨ।

branch replay ਕੰਮ ਦੇ ਬੋਝ ਨੂੰ ਕਿਵੇਂ ਘਟਾਉਂਦਾ ਹੈ

LongStraw ਦਾ ਤਰੀਕਾ transformer generation ਬਾਰੇ ਦੋ ਨਿਰੀਖਣਾਂ 'ਤੇ ਅਧਾਰਤ ਹੈ:

  • Rollouts ਵਿੱਚ ਪ੍ਰੋਂਪਟ (prompt) ਅਤੇ ਇੱਕ ਜਵਾਬ ਦਾ ਸ਼ੁਰੂਆਤੀ ਹਿੱਸਾ ਇੱਕੋ ਜਿਹਾ ਹੁੰਦਾ ਹੈ।
  • ਹਰੇਕ ਜਵਾਬ ਦੇ ਸਿਰਫ਼ ਵੱਖਰੇ (divergent) ਪਿੱਛਲੇ ਹਿੱਸੇ (tail) ਨੂੰ ਹੀ ਅਸਲ ਵਿੱਚ ਨਵੇਂ computation ਦੀ ਲੋੜ ਹੁੰਦੀ ਹੈ।

ਇਹ ਸਿਸਟਮ ਇੱਕ architecture-aware execution stack ਬਣਾਉਂਦਾ ਹੈ ਜੋ shared prefix ਲਈ activations ਨੂੰ ਰਿਕਾਰਡ ਕਰਦਾ ਹੈ। ਜਦੋਂ ਇੱਕ ਨਵੀਂ branch ਦੀ ਜਾਂਚ ਕੀਤੀ ਜਾਂਦੀ ਹੈ, ਤਾਂ ਇਹ ਉਸ ਨੂੰ ਦੁਬਾਰਾ ਕੰਪਿਊਟ ਕਰਨ ਦੀ ਬਜਾਏ cached prefix ਨੂੰ replay ਕਰਦਾ ਹੈ, ਅਤੇ ਫਿਰ ਸਿਰਫ਼ ਨਵੇਂ ਹਿੱਸੇ (novel segment) 'ਤੇ ਹੀ backward pass ਚਲਾਉਂਦਾ ਹੈ। ਅਸਲ ਵਿੱਚ ਇਸਦਾ ਮਤਲਬ ਹੈ ਕਿ backward pass ਬਹੁਤ ਘੱਟ token positions ਨੂੰ ਛੂਹਦਾ ਹੈ, ਜਿਸ ਨਾਲ raw compute ਵਿੱਚ 8 ਤੋਂ 16 ਗੁਣਾ ਦੀ ਕਮੀ ਆਉਂਦੀ ਹੈ।

ਤਤਕਾਲ ਪ੍ਰਭਾਵ

  • ਅੱਠ H20 GPUs 'ਤੇ 2.1 M token positions ਪ੍ਰੋਸੈਸ ਕੀਤੀਆਂ ਗਈਆਂ, ਇੱਕ ਅਜਿਹਾ ਹਾਰਡਵੇਅਰ ਬਜਟ ਜੋ ਆਮ ਤੌਰ 'ਤੇ ਉਸ ਕੰਮ ਦੇ ਬੋਝ ਦੇ ਇੱਕ ਛੋਟੇ ਹਿੱਸੇ ਦਾ ਹੀ ਸਮਰਥਨ ਕਰ ਸਕਦਾ ਸੀ।
  • Long-context RL ਵਿੱਚ ਰੁਕਾਵਟ (bottleneck) ਨੂੰ ਸਿੱਧਾ ਨਿਸ਼ਾਨਾ ਬਣਾਉਣਾ, ਜਿੱਥੇ context ਵਧਣ ਨਾਲ memory ਅਤੇ compute costs ਬਹੁਤ ਜ਼ਿਆਦਾ ਵਧ ਜਾਂਦੇ ਹਨ।
  • ਲੈਬਾਂ GPU ਅਲਾਟੇਸ਼ਨ ਨੂੰ ਬਦਲ ਸਕਦੀਆਂ ਹਨ: ਉਹੀ ਹਾਰਡਵੇਅਰ, ਜੋ ਮੁੱਖ ਤੌਰ 'ਤੇ ਇੱਕ inference accelerator ਹੈ, ਹੁਣ ਟ੍ਰੇਨਿੰਗ ਲਈ ਵਰਤਿਆ ਜਾ ਸਕਦਾ ਹੈ, ਹਾਲਾਂਕਿ ਹੋਰ ਕਾਰਡਾਂ 'ਤੇ ਨਤੀਜੇ ਵੱਖਰੇ ਹੋ ਸਕਦੇ ਹਨ।

ਖੁੱਲ੍ਹੇ ਸਵਾਲ ਅਤੇ ਸੀਮਾਵਾਂ

ਐਲਾਨ ਵਿੱਚ ਟ੍ਰੇਨਿੰਗ ਦੀ ਰਫ਼ਤਾਰ ਦੇ ਅੰਕੜੇ ਅਤੇ convergence curves ਨਹੀਂ ਦਿੱਤੇ ਗਏ ਹਨ, ਇਸ ਲਈ ਸਾਨੂੰ ਇਹ ਨਹੀਂ ਪਤਾ ਕਿ compute ਵਿੱਚ ਕਟੌਤੀ ਦਾ ਮਤਲਬ ਤੇਜ਼ wall-clock ਸਮਾਂ ਹੈ ਜਾਂ ਸਿਰਫ਼ ਘੱਟ GPU occupancy ਹੈ। ਇਹ ਵਿਧੀ autoregressive sampling ਲਈ ਦੱਸੀ ਗਈ ਹੈ; non-autoregressive ਜਾਂ hybrid ਰਣਨੀਤੀਆਂ ਨਾਲ ਇਸਦਾ ਵਿਵਹਾਰ ਅਜੇ ਅਜ਼ਮਾਇਆ ਨਹੀਂ ਗਿਆ ਹੈ। ਕਿਉਂਕਿ H20 ਮੁੱਖ ਤੌਰ 'ਤੇ ਇੱਕ inference accelerator ਹੈ, ਇਸ ਲਈ H100 ਜਾਂ B200 ਵਰਗੇ ਵਧੇਰੇ ਆਮ ਟ੍ਰੇਨਿੰਗ ਕਾਰਡਾਂ 'ਤੇ ਪ੍ਰਦਰਸ਼ਨ ਵੱਖਰਾ ਹੋ ਸਕਦਾ ਹੈ।

ਸੁਤੰਤਰ benchmarks ਨੇ ਅਜੇ ਤੱਕ LongStraw ਦੇ ਅੰਕੜਿਆਂ ਦੀ ਪੁਸ਼ਟੀ ਨਹੀਂ ਕੀਤੀ ਹੈ। Third-party validation ਤੋਂ ਬਿਨਾਂ, ਭਾਈਚਾਰੇ ਨੂੰ ਇਹਨਾਂ ਨਤੀਜਿਆਂ ਨੂੰ ਉਮੀਦ ਭਰਪੂਰ ਪਰ ਅਸਥਾਈ (provisional) ਮੰਨਣਾ ਚਾਹੀਦਾ ਹੈ।

ਕੀ ਦਾਅ 'ਤੇ ਹੈ

ਜੇਕਰ branch-replay ਦਾ ਵਿਚਾਰ ਹੋਰ RL fine-tuning algorithms ਜਿਵੇਂ ਕਿ Direct Preference Optimization (DPO) ਜਾਂ Proximal Policy Optimization (PPO) ਤੱਕ ਫੈਲਦਾ ਹੈ, ਤਾਂ long-context ਮਾਡਲਾਂ ਲਈ ਖਰਚੇ ਦੀ ਰੁਕਾਵਟ ਖਤਮ ਹੋ ਸਕਦੀ ਹੈ।

ਕੀ ਦੇਖਣਾ ਚਾਹੀਦਾ ਹੈ

  • GPU architectures ਦੀ ਇੱਕ ਲੜੀ 'ਤੇ third-party replication ਦੇ ਯਤਨ।
  • Baseline pipelines ਦੇ ਮੁਕਾਬਲੇ training throughput ਅਤੇ ਅੰਤਿਮ ਮਾਡਲ ਦੀ ਗੁਣਵੱਤਾ 'ਤੇ LongStraw ਤੋਂ ਅਪਡੇਟਸ।