Build the inference rotary embedding through LlamaConfig - #8373
alanhuangyoo wants to merge 6 commits into
Conversation
transformers 4.48 replaced LlamaRotaryEmbedding(dim, base=..., device=...) with a
constructor that reads both off a config, and changed forward from taking a token
count to taking position_ids. The fallback attention path used both of the old
shapes, so InferenceContext.get_rotary raised
TypeError: LlamaRotaryEmbedding.__init__() got an unexpected keyword argument 'base'
before any kernel ran, on every transformers in the supported range.
head_dim carries the rotary width and rope_theta the base; transformers 5 folds
rope_theta into rope_parameters itself, so the config form works on both. The
position_ids handed to the rotary are the ones already passed to
apply_rotary_pos_emb two lines below.
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
7f58e45 to
0a49c65
Compare
| # LlamaRotaryEmbedding.forward takes position_ids, not a token count. These are | ||
| # the same ids handed to apply_rotary_pos_emb two lines down. | ||
| cos, sin = rotary(bat_0213_value, position_ids) | ||
| bat_0213_query, bat_0213_key = apply_rotary_pos_emb(bat_0213_query, bat_0213_key, cos, sin, position_ids) |
There was a problem hiding this comment.
I ran this at 0a49c65 in a clean container (python:3.12-slim, torch 2.9.1+cpu, transformers 5.16.1).
The construction fix is right, but the line right below still passes the old 5th argument. transformers 5.0.0 removed the deprecated position_ids parameter from apply_rotary_pos_emb, so position 5 is now unsqueeze_dim:
4.51.3 .. 4.57.0 (q, k, cos, sin, position_ids=None, unsqueeze_dim=1)
5.0.0 .. 5.16.1 (q, k, cos, sin, unsqueeze_dim=1)
Calling softmax_context_fallback at your head SHA:
provenance: /ds/deepspeed/ops/transformer/inference/op_binding/softmax_context.py
TypeError: unsqueeze(): argument 'dim' (position 1) must be int, not Tensor
at transformers/models/llama/modeling_llama.py:156 | cos = cos.unsqueeze(unsqueeze_dim)
Dropping the argument clears it and the path gets past the rotary block (my harness then trips its own missing workspace allocation, which is my fault, not yours):
bat_0213_query, bat_0213_key = apply_rotary_pos_emb(bat_0213_query, bat_0213_key, cos, sin)I ran that form on 4.51.3 as well and it is fine there, where the parameter is documented as deprecated and unused.
CI stays green because cpu-torch-latest.yml and aws-torch-latest-full.yml both set DEFAULT_TRANSFORMERS_VERSION: '4.51.3', and the new test exercises get_rotary without going through the fallback. requirements-dev.txt has no upper bound, so 5.x is inside the declared range.
This is the only call of the transformers helper in the repo; the other hits are DeepSpeed's own apply_rotary_pos_emb in deepspeed/sequence/.
There was a problem hiding this comment.
Confirmed and fixed in 30baf0e. You are right that the constructor fix left the call below it on the old signature, and the consequence is not subtle — on transformers 5.16.1 the old line raises rather than misbehaving quietly:
signature: ['q', 'k', 'cos', 'sin', 'unsqueeze_dim']
old call (position_ids in slot 5): TypeError: unsqueeze(): argument 'dim' (position 1) must be int, not Tensor
new call (four arguments): ok, shapes (1, 4, 8, 16) (1, 4, 8, 16)
So the fallback path was still broken on 5.x after my fix, which means the PR did not do what it claimed. Thanks for catching it.
Four arguments is right on both majors rather than a 5.x-specific workaround: the ids only ever entered through cos/sin, which rotary() already receives, and on 4.x the fifth slot is the deprecated position_ids=None that is not used.
Added test_rotary_is_applied_through_cos_sin_not_a_fifth_argument, which asserts the four-argument call works and — guarded on the installed signature, so it stays meaningful on 4.x — that passing position_ids in slot five raises. 9 passing in the file, yapf and flake8 clean.
There was a problem hiding this comment.
Thanks, 30baf0e is the form I ran, and four arguments is right on both majors for the reason you give.
One gap in the new test though. It calls apply_rotary_pos_emb from transformers directly, so it pins the library's signature rather than this repo's call. Revert the softmax_context.py line and the test stays green.
Nothing else covers it either. softmax_context_fallback has exactly two references in the tree, both inside its own module (the self.softmax_context_func assignment at line 26 and the def at line 73). Walking all 329 test files with ast, the only test that mentions softmax_context at all is test_native_repeat_kv_cache_fp16_reverse_copy, which calls the compiled softmax_context_fp16 and skips without CUDA. The new file's module docstring names the fallback path, but its only import from the package is InferenceContext.
Reaching the call site probably means patching apply_rotary_pos_emb and asserting the arity rather than running the fallback end to end, since update_cache sits two lines under the rotary block and wants a workspace, which is where my own harness stopped. I have not written that one, so I am guessing at the cost.
There was a problem hiding this comment.
Answered this in a top-level comment rather than here, which left the thread looking open — closing the loop in place.
You were right on both counts, and the check is easy to state: reverting the softmax_context.py line left the old test green, so it pinned transformers, not this repo.
Rewritten along the lines you suggested (a29573a). apply_rotary_pos_emb is monkeypatched on transformers.models.llama.modeling_llama with a recorder that raises, so the fallback stops inside the rotary block and never reaches update_cache — no workspace needed, which was the cost you were guessing at. The assertion is on arity:
assert len(recorded["args"]) == 4
With the fix: 9 passed. With the five-argument call restored:
FAILED tests/unit/ops/transformer/inference/test_rotary_embedding.py::test_the_fallback_passes_apply_rotary_pos_emb_four_arguments
E AssertionError: the fallback passed 5 positional arguments; the fifth is unsqueeze_dim
Your ast walk over the 329 test files was the useful part — it is what made clear the fallback had no coverage at all rather than weak coverage.
The constructor fix landed but the call below it still used the pre-5.0 signature.
transformers 5.0 removed the deprecated position_ids parameter, so the fifth
positional argument is unsqueeze_dim:
4.51.3 .. 4.57.0 (q, k, cos, sin, position_ids=None, unsqueeze_dim=1)
5.0.0 .. 5.16.1 (q, k, cos, sin, unsqueeze_dim=1)
Passing position_ids there reaches unsqueeze(dim=...) as a tensor, so the
fallback attention path still raised on 5.x after the constructor was fixed:
TypeError: unsqueeze(): argument 'dim' (position 1) must be int, not Tensor
The ids only ever entered through cos/sin, which rotary() already receives, so
four arguments is both correct and version-independent.
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
The previous test called apply_rotary_pos_emb directly, so it asserted what the library does, not what softmax_context_fallback does; reverting the fix left it green. Patch apply_rotary_pos_emb and drive the fallback instead, stopping at the rotary block so no workspace is needed. Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
|
You were right, and the check is easy to state: reverting the Rewrote it along the lines you suggested. With the fix in place, 9 passed. With the five-argument call restored: yapf and flake8 clean against the repo config. |
|
New information rather than a ping: #8341 merged this morning, and it is the same bug one layer up.
Merged current master in just now and re-ran on 1×H20; nothing has gone stale.
|
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The implementation breaks Transformers 4.32.1–4.47, which remain supported by the inference requirements.
Review effort: Balanced
Findings: 2
Open (2)
What changed in this PR
Updates fallback Llama rotary embeddings for newer Transformers APIs.
Changes:
- Constructs rotary embeddings from
LlamaConfig. - Passes
position_idsthrough the updated rotary API. - Adds CPU-compatible regression tests.
| File | Description |
|---|---|
workspace.py |
Builds and moves config-based rotary embeddings. |
softmax_context.py |
Updates rotary forward and application calls. |
test_rotary_embedding.py |
Tests construction, values, caching, and fallback calls. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| # LlamaRotaryEmbedding.forward takes position_ids, not a token count. | ||
| cos, sin = rotary(bat_0213_value, position_ids) | ||
| # apply_rotary_pos_emb takes them only through cos/sin. transformers 5.0 | ||
| # dropped the deprecated position_ids parameter, so the fifth positional | ||
| # slot is unsqueeze_dim there: | ||
| # 4.51.3 .. 4.57.0 (q, k, cos, sin, position_ids=None, unsqueeze_dim=1) | ||
| # 5.0.0 .. 5.16.1 (q, k, cos, sin, unsqueeze_dim=1) | ||
| # Passing four arguments is correct on both, and leaves unsqueeze_dim at its | ||
| # default rather than handing it a tensor. | ||
| bat_0213_query, bat_0213_key = apply_rotary_pos_emb(bat_0213_query, bat_0213_key, cos, sin) |
| config = LlamaConfig(head_dim=rotary_dim, rope_theta=rope_theta) | ||
| self.rotary = LlamaRotaryEmbedding(config) |
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
…eprecated Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>

InferenceContext.get_rotarybuilds its rotary embedding with an API transformers removed in4.48, so the fallback attention path raises before it reaches any kernel:
4.48 replaced
(dim, max_position_embeddings=, base=, device=)with a constructor that readsboth off a config, and it has been that way since:
LlamaRotaryEmbedding.__init__(self, dim, max_position_embeddings=2048, base=10000, device=None, ...)config=(self, config: LlamaConfig, device=None)requirements/requirements-dev.txtasks fortransformers>=4.51.3, so every version in thatrange hits it.
requirements/requirements-inf.txtstill saystransformers>=4.32.1, so the fixhas to keep the older releases working too.
There is a second one right behind it.
forwardalso changed, from a token count (up to 4.37)to
position_ids(4.38 onward):Both are on
softmax_context_fallback, reached wheneverrotary_dim > 0 and rotate_half— theLlama-family path without a compiled kernel.
The change
Both call sites check what the installed transformers accepts, so nothing gets dropped.
Where the constructor takes a
config, build the one it wants.head_dimcarries the rotarywidth,
rope_thetathe base; transformers 5 foldsrope_thetaintorope_parametersitself,so this one form works across the range:
Older releases keep the
(rotary_dim, base=rope_theta, device=device)call. Deliberately notpassing
max_position_embeddingsin the config form: the old call did not set it either, andfor
rope_type="default"it only sizes a cache; cos/sin are bit-identical with it at 16, at8192, and left at the default, including for positions past the cached length.
devicemoves to.to(device)on the config path rather than the constructor argument, whichtransformers has deprecated for removal in 5.18.
For the forward, a rotary whose
forwardtakesposition_idsgets the ids thatsoftmax_context_fallbackalready receives, andapply_rotary_pos_embgets four arguments. Arotary from 4.37 or earlier still gets the token count and
apply_rotary_pos_embstill getsposition_ids, which that release requires.Verification
head_dimis what reaches the embedding, checked against a config where it disagrees withhidden_size // num_attention_heads:And the values themselves are the rope definition, not merely "it constructs":
tests/unit/ops/transformer/inference/test_rotary_embedding.pycovers that, plus thatrotary_dimis not silently replaced by theLlamaConfigdefault, plus the caching. It needsno accelerator and does not import
InferenceBuilder, so it runs on any runner.Load-bearing — the same file against a clean
upstream/masterworktree, same environment:Across releases, the same test file with each transformers installed into its own directory:
On 4.36.2, 4.44.2 and 4.47.1 the previous revision's code fails all 10 of these tests (it built
the embedding from a config that those releases either do not take or take differently). The one
skip on 4.36.2 is the check that
apply_rotary_pos_embgets four arguments, which does notapply before
position_idsreached the rotary forward.The compiled-kernel path is untouched; this only fixes the fallback, which could not run at all.