-
Notifications
You must be signed in to change notification settings - Fork 1.1k
[MLX] Add off-graph KV cache ring runtime #21532
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
kiymetakdemir
wants to merge
4
commits into
pytorch:main
Choose a base branch
from
kiymetakdemir:kv-cache-mlx-ring
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+253
−35
Open
Changes from all commits
Commits
Show all changes
4 commits
Select commit
Hold shift + click to select a range
c921c59
[MLX] Add ring layers with a declarative sliding-window mask
kiymetakdemir 7249708
[MLX] Build the sliding-window mask in the cache
kiymetakdemir c39bb1f
Merge branch 'main' into kv-cache-mlx-ring
kiymetakdemir 94187b1
[MLX] Drop a duplicated comment in the attend handler
kiymetakdemir File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Creating this mask is not free, are you creating it once per layer IIUC
Per decode step, we only need to plan once per policy per execute step. Is that how your code is set up?
Much of this will become more apparent when we test in real model.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Decode currently doesn't need a mask. But no, we don't plan once per policy. Policies aren't duplicated, but update_and_fetch is per layer, so we call plan() for each layer per step. I can memoize the plan and mask per policy index, invalidated when (position, T) changes.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Why doesn't decode need a mask? Isn't the mask how you tell the ring what to attend to?
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The selection happens in the read, the planner returns the retained window as physical runs, and we read exactly those. For example, W=4, max_write=1 (ring of 4 slots), decoding at position 10. RingPolicy::plan(10, 1) gives rstart = 7, span 4, which wraps into two runs, {start:3, len:1} and {start:0, len:3}. We read those and concatenate, resulting positions 7, 8, 9, 10 in logical order.
We need a mask during prefill: a step with T > 1 reads the union of its queries' visibilities.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Even in prefill, I'm not sure we'll want to generate a mask per layer.
Stamping to unblock, but consider doing plan once per execute per policy, not per layer.
I guess the next PR will be E2E enablement, and we can start comparing perf.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Thanks! Do you think memoization would make sense for this problem? Keeping the plan (and mask) per policy index and invalidating when (position, T)?