Skip to content

Add support for YaRN scaling - #477

Merged
jonatanklosko merged 4 commits into
elixir-nx:mainfrom
dkuku:feat/rope-yarn
Oct 5, 2026
Merged

jonatanklosko merged 4 commits into
elixir-nx:mainfrom
dkuku:feat/rope-yarn

Conversation

@dkuku

@dkuku dkuku commented Oct 4, 2026 •

Copy link
Copy Markdown
Contributor

Pr includes #476

This pr ads native YaRN scaling support across models that use rotary embeddings:

• Frequency blending between extrapolation and interpolation using a linear ramp over beta_fast and beta_slow.
• Attention scale adjustment (attention_factor / mscale).
• Automatic parsing from Hugging Face rope_scaling / rope_parameters in Bumblebee.Shared.

It is used by models such as NousResearch/Yarn-Llama-2-7b-64k, NousResearch/Yarn-Llama-2-7b-128k, or NousResearch/Yarn-Mistral-7b-128k

Reproduction Steps

Switched to branch 'main'
Your branch is ahead of 'origin/main' by 2 commits.
  (use "git push" to publish your local commits)
❯ iex -S mix
Erlang/OTP 29 [erts-17.0] [source] [64-bit] [smp:16:16] [ds:16:16:10] [async-threads:1] [jit:ns]

Compiling 59 files (.ex)
Interactive Elixir (1.20.0) - press Ctrl+C to exit (type h() ENTER for help)
  iex(1)> {:ok, spec} = Bumblebee.load_spec({:hf, "NousResearch/Yarn-Llama-2-7b-64k"})
** (RuntimeError) conversion failed, unsupported rotary embedding parameters: %{"factor" => 16.0, "finetuned" => true, "original_max_position_embeddings" => 4096, "type" => "yarn"}
    (bumblebee 0.8.0) lib/bumblebee/shared.ex:215: Bumblebee.Shared.rotary_embedding_scaling_strategy/2
    (bumblebee 0.8.0) lib/bumblebee/shared.ex:174: Bumblebee.Shared.rotary_embedding_options/2
    (bumblebee 0.8.0) lib/bumblebee/text/llama.ex:412: Bumblebee.HuggingFace.Transformers.Config.Bumblebee.Text.Llama.load/2
    (bumblebee 0.8.0) lib/bumblebee.ex:489: Bumblebee.do_load_spec/4
    iex:1: (file)
❯ gco -
Switched to branch 'feat/rope-yarn'
Your branch is up to date with 'origin/feat/rope-yarn'.
❯ iex -S mix
Erlang/OTP 29 [erts-17.0] [source] [64-bit] [smp:16:16] [ds:16:16:10] [async-threads:1] [jit:ns]

Compiling 59 files (.ex)
Interactive Elixir (1.20.0) - press Ctrl+C to exit (type h() ENTER for help)
iex(1)> {:ok, spec} = Bumblebee.load_spec({:hf, "NousResearch/Yarn-Llama-2-7b-64k"})
{:ok,
 %Bumblebee.Text.Llama{
   architecture: :for_causal_language_modeling,
   vocab_size: 32000,
   max_positions: 65536,
...

Codex and Antigravity was used to implement this feature

@dkuku

dkuku commented Oct 4, 2026 •

Copy link
Copy Markdown
Contributor Author

python results:

uv run --with transformers --with torch python test.py
Loading weights: 100%|█████████████| 20/20 [00:00<00:00, 43599.83it/s]
Python hidden_state slice:
[[[1.480570673942566, -2.031874179840088, 0.4768206477165222], [2.373222589492798, -0.8378004431724548, -0.022018365561962128], [0.5740560293197632, -0.05190371721982956, -1.180698275566101]]]
import torch
from transformers import LlamaModel
from transformers.models.llama.modeling_llama import LlamaRotaryEmbedding

model = LlamaModel.from_pretrained('bumblebee-testing/tiny-random-LlamaModel')

model.config.rope_scaling = {
    'rope_type': 'yarn',
    'rope_theta': 10000.0,
    'factor': 4.0,
    'original_max_position_embeddings': 16,
    'beta_fast': 32.0,
    'beta_slow': 1.0,
    'attention_factor': 1.138629436111989
}
model.rotary_emb = LlamaRotaryEmbedding(model.config)

input_ids = torch.tensor([[10, 20, 30, 40, 50, 60, 70, 80, 0, 0]])
attention_mask = torch.tensor([[1, 1, 1, 1, 1, 1, 1, 1, 0, 0]])

with torch.no_grad():
    out = model(input_ids=input_ids, attention_mask=attention_mask)

print("Python hidden_state slice:")
print(out.last_hidden_state[:, 1:4, 1:4].tolist())

@jonatanklosko jonatanklosko changed the title Feat/rope yarn Add support for YaRN scaling Oct 5, 2026
@jonatanklosko
jonatanklosko merged commit dc3150c into elixir-nx:main Oct 5, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants