Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright © 2023 Apple Inc.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,201 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -0,0 +1,41 @@
|
||||
MTPLX
|
||||
Copyright 2026 Youssof Altoukhi
|
||||
|
||||
MTPLX is a native MTP speculative decoding project for Apple Silicon.
|
||||
|
||||
ATTRIBUTION REQUIREMENT
|
||||
|
||||
This NOTICE file is part of the Apache License 2.0 terms for MTPLX (see
|
||||
section 4(d) of the LICENSE). Any product, application, service, or
|
||||
distribution that includes, embeds, or is built on MTPLX, in whole or in
|
||||
part, modified or unmodified, must display the following attribution within
|
||||
the product itself, in a place a user of that product can see (for example an
|
||||
About screen, a credits or acknowledgements screen, a settings or help page,
|
||||
documentation shipped with the product, or the startup banner of a command
|
||||
line tool):
|
||||
|
||||
Powered by MTPLX
|
||||
https://github.com/youssofal/mtplx
|
||||
|
||||
Attribution in a source repository, a README, or a marketing page alone does
|
||||
not satisfy this requirement. The words "Powered by MTPLX" must appear
|
||||
in-product. The link is required wherever the display medium supports it.
|
||||
|
||||
Public benchmarks, articles, and research that use or build on MTPLX should
|
||||
credit "MTPLX by Youssof Altoukhi" with the same link.
|
||||
|
||||
If MTPLX informs academic or technical writing, please cite the repository using
|
||||
the included CITATION.cff metadata.
|
||||
|
||||
This distribution includes a vendored standalone subset of vllm-metal's
|
||||
Apache-2.0 licensed Metal paged-attention kernels under vllm_metal/metal.
|
||||
The vendored subset is used only for local MLX/Metal kernel dispatch and does
|
||||
not include or depend on the vLLM serving stack.
|
||||
|
||||
This product includes Metal kernel code adapted from dflash-mlx
|
||||
(https://github.com/bstnxbt/dflash-mlx), Copyright dflash-mlx contributors,
|
||||
licensed under the Apache License 2.0. See mtplx/nax_verify.py for details.
|
||||
|
||||
This product includes the Apache-2.0 licensed MLX implementation for
|
||||
Laguna-S-2.1 from PipeNetwork, Copyright 2026 PipeNetwork, under
|
||||
mtplx/models/laguna.py.
|
||||
@@ -0,0 +1,551 @@
|
||||
MTPLX Qwen runtime shaders
|
||||
=========================
|
||||
|
||||
mtplx-runtime-0.32.2.metallib is an unchanged copy of mlx/lib/mlx.metallib
|
||||
from the existing MTPLX reference environment (runtime version 0.32.2).
|
||||
|
||||
SHA256: dc59d1cceb1a5c7e578232e6e41e28e2c73c9463ac6dbc3886c3ee17ffc270ed
|
||||
Source tag: v0.32.2
|
||||
Source commit: 1f8e74e3f12f31365464a6867c6579f0e9b29d85
|
||||
Source: https://github.com/ml-explore/mlx/tree/v0.32.2
|
||||
License: MIT, reproduced in MLX-LM-LICENSE.txt (identical license text).
|
||||
|
||||
Only GPU shader code is included. No libmlx.dylib, Python inference code,
|
||||
or C/C++ application/backend code is linked into DS4Server. Rust selects
|
||||
the shader entry points, binds model buffers and owns graph execution.
|
||||
|
||||
The reference's get_quantized_kernel implementation was verified to call
|
||||
Device::get_kernel on the default precompiled library, matching
|
||||
mlx/backend/metal/nojit_kernels.cpp at the pinned commit. The shader library
|
||||
is reused byte-for-byte, including its precompiled specializations, rather
|
||||
than independently rebuilding different kernels or compiler settings.
|
||||
|
||||
Provenance on the evaluation machine:
|
||||
/private/tmp/MTPLX-analysis-20260901/.venv/lib/python3.12/site-packages/mlx/lib/mlx.metallib
|
||||
|
||||
The Rust test mtplx_runtime_shaders_are_pinned enforces the complete file
|
||||
hash. Existing resource packaging includes the metal directory recursively.
|
||||
This artifact adds approximately 174 MiB to the resources; it is not a model
|
||||
download. Loading is lazy on the first runtime-shader dispatch.
|
||||
|
||||
The shader identity is not a claim that the complete Rust model graph,
|
||||
scheduling, cache behavior, or end-to-end performance has reached parity.
|
||||
|
||||
Compiled QSA moving offsets
|
||||
--------------------------
|
||||
|
||||
tests/fixtures/mtplx-qsa-update-jit.json records the actual compiled indexer
|
||||
scalar shaders and compute_dynamic_offset_int32, observed without changing
|
||||
the original compilation by tools/mtplx-jit-reference.py --operation qsa-update.
|
||||
Only the contiguous int32[1] scalar specializations are retained; full observed
|
||||
library source hashes are included. --check reruns the original compiled core.
|
||||
The Dynamic Offset body is used unchanged by the Rust dynamic-copy path;
|
||||
its original full source hash is
|
||||
48a7309664f797e749aa42d2c2c4db0cf3abedf97297f0068d06b7847d988b93.
|
||||
It is Copyright Apple Inc., MIT as reproduced in MLX-LM-LICENSE.txt.
|
||||
The following gg1/gg2_dynamic_copybfloat16bfloat16 kernels are taken directly
|
||||
from the unchanged runtime metallib. No frontiers are read back to the CPU.
|
||||
Clamp/Multiply now feed the connected qsa_compiled_cache_window stage.
|
||||
mtplx-qsa-compiled-scalars.json retains the actual generated scalar kernels,
|
||||
including all three constant-CSE layouts for the Clamp and an independent
|
||||
257/255 specialization check. Only structural integer literals and exported
|
||||
symbols change when Rust specializes a kernel; the computations are unchanged.
|
||||
mtplx-qsa-compiled-header.metal is the unmodified original compiler prefix;
|
||||
SHA256: 2665a76463f3f6ee283c6a50b66e4a527318a114080b31441dfa900042097a39.
|
||||
--qsa-header --check compares that prefix with a fresh original compilation.
|
||||
Both resources contain runtime shader code, not host runtime code. Most is
|
||||
Apple MIT. The unchanged full prefix also retains the Apache-2.0 cexpf.h
|
||||
notice (Apple, NVIDIA, Filipe RNC Maia; license text in MTPLX-LICENSE.txt) and
|
||||
the full BSD-2-Clause expm1f.h notice/disclaimer (Norbert Juffa 2015-2023).
|
||||
Those overloads are retained as original header dependencies, not new Qwen
|
||||
complex/exponential computation in the integer scalar kernels.
|
||||
26 compiled reference calls verify the complete retained-input cache window.
|
||||
They do not establish graph-bank replay, allocation/donation, BFS scheduling
|
||||
or production inference parity.
|
||||
|
||||
tests/fixtures/mtplx-qsa-select-jit.json also retains the actual compiled Add
|
||||
kernels for selector row offsets 0/1/2/3. Rust substitutes only the structural
|
||||
offset literal and exported symbol, preserving the original +0 dispatch.
|
||||
--operation qsa-update --qsa-mode blocks --qsa-score-budget 4096 --qsa-header
|
||||
--check reproduces the connected reference selector and its exact sources.
|
||||
Query preparation and both selector families now accept GPU frontier leaves
|
||||
through the same original kernel dispatch used by the host-frontier entry.
|
||||
66 actual compiled reference calls cover all five output modes, chunked
|
||||
selection and the connected cache state for retained old input leaves. These
|
||||
are functional Q/K-entry checks, not graph-bank, ownership or UI performance
|
||||
acceptance. Host integration and the runtime evaluator remain open.
|
||||
The retained-input Hidden entry has since been connected through the already
|
||||
verified original affine projection kernels to that same Q/K implementation.
|
||||
132 actual select_hidden calls cover 4/8-bit, group32/64 projections and all
|
||||
five output modes. The combined entry test covers 198 reference calls. No new
|
||||
Metal bodies or alternative projection/selection arithmetic are introduced.
|
||||
The installed B1/BF16 cache/phase routing is now connected to backing reserve,
|
||||
explicit GPU frontiers, those same retained-input arithmetic entries and cache
|
||||
commit. 720 actual original host-method decisions and 108 additional complete
|
||||
indexer calls cover routing and ongoing state/output transitions. The combined
|
||||
host-flow test includes the previous 92 non-compiled calls through the same
|
||||
entry. A parameter-bound QSA graph bank now replaces the direct compiled
|
||||
expression chains. It rebinds explicit inputs to cached primitive dependencies
|
||||
and uses the pinned degree/BFS-width algorithm for a single indexer graph.
|
||||
99 optimized original graph contracts, all 198 arithmetic calls, stride-changing
|
||||
replay, parameter invalidation and the original connected-call trace/entry
|
||||
counters pass. Evaluated constants are omitted from structural fingerprints;
|
||||
kernel-source and output checks remain separate. No alternate Metal kernel was
|
||||
introduced. The dtype/shape-generic guard, donation/allocator, early release,
|
||||
global model scheduling and production integration remain open. Last-use
|
||||
graph leaves are now detached after their consumer, separately from explicit
|
||||
completion ownership that protects GPU work until its existing CB finishes.
|
||||
The shared canonical dispatch bridge holds bound Metal resources through
|
||||
completion as well; this fixes four Invalid Resource failures exposed with
|
||||
unretained command buffers. Compile/dispatch use scoped autorelease pools.
|
||||
The normal and strengthened model-free collections cover 34 tests. This is
|
||||
not performance acceptance or evidence of matching whole-model Metal encoder
|
||||
timelines. Donation must still account for both descriptor and shared Data
|
||||
ownership, including outstanding GPU evaluator holds.
|
||||
The canonical test-bound Buffer now separates array/view identity from shared
|
||||
Data ownership, including nested native views and completion holds. QSA COW
|
||||
checks both Rc<Buffer> sharing and underlying Data sharing. A direct pinned
|
||||
QSACache alias/view update and its negative Rust regression check prove that
|
||||
array aliases observe replacement while distinct views retain old values.
|
||||
The QSA graph now applies the pinned primitive input/sibling-minus-primary-output
|
||||
Data retention protocol before the next primitive, rather than re-holding leaves
|
||||
at their last consumer. Its scheduler tracks the actually selected sibling, not
|
||||
just the producer node. CPU scheduling and GPU ownership-count checks cover that
|
||||
distinction, duplicate Data, empty-batch fallback and completion.
|
||||
QSA raw/pool DynamicSliceUpdate now performs actual BF16 vector-copy donation
|
||||
for exclusive mutable cache state. Retained inputs, snapshots, views and GPU
|
||||
Data holds select copying instead. The 16 KiB bound uses root allocation size.
|
||||
All 198 core reference cases also run with state snapshots and exclusive state,
|
||||
checking old/new hashes and actual Data reuse. No kernels or synchronization
|
||||
boundaries changed. Model-wide integration, other primitive donation and generic
|
||||
dtype/layout contracts remain open.
|
||||
|
||||
The allocator policy is now ported from the pinned buffer_cache.h with the same
|
||||
best-fit/oldest-equal-size choice, strict reuse ceiling and age-based/90%-clear
|
||||
eviction. A 101-event trace from the real installed runtime checks allocation
|
||||
identities and active/cached bytes, including cache-limit transitions. Reference page rounding and the
|
||||
device maxBufferLength precheck are connected to the test-bound Buffer methods;
|
||||
logical views preserve tensor bounds while Data records the rounded root size.
|
||||
The explicit Rust Allocator now owns a real 1 MiB untracked/shared Metal heap,
|
||||
uses it for requests below 256 bytes with device-allocation fallback, and
|
||||
recycles native roots only after the final physical allocation hold releases.
|
||||
Its active/cache/peak/resource accounting, cache and memory limits, resource
|
||||
pressure GC (including the original unsigned subtraction), zero/null result
|
||||
and actual cached storage reuse are checked against the 101-event receipt.
|
||||
Residency/wired limits are now connected to this explicit allocator, including
|
||||
heap registration, cache retention and erase-before-release. Set selection and
|
||||
budgets follow resident.cpp: first fit, oversize/empty-set reuse, 32-set ceiling,
|
||||
emptiest fallback and touched-set commits on resize. Native membership and the
|
||||
ten original residency lifecycle scenarios are tested. Queue attachment uses
|
||||
the published set count and is exercised immediately before test commits;
|
||||
The new test-bound Submission owner now attaches automatically at actual native
|
||||
commit boundaries through a scoped encoding-thread callback, including flush,
|
||||
readback, finish and cleanup paths. Its queue cursor persists across batches;
|
||||
scope teardown drains before unregistering the callback. Legacy work outside
|
||||
the scope and other queues do not inherit it. Externally wrapped storage,
|
||||
the process-wide owner and model-wide routing remain open.
|
||||
Existing canonical helpers
|
||||
are not globally switched to untracked buffers before encoder dependencies
|
||||
are ported. No whole-model allocator/performance parity is claimed.
|
||||
|
||||
The test-bound Rust Encoder now owns an independent queue with unretained
|
||||
command buffers and Concurrent compute encoders, following the pinned
|
||||
device.cpp/event.cpp/error.h dependency and completion rules. Access roles,
|
||||
barrier epochs, deferred concurrent outputs, cross-encoder fences, temporary
|
||||
exclusion and shared-event error propagation are managed in Rust. The bridge
|
||||
only issues Metal API calls and reuses the same original kernel dispatcher.
|
||||
Commit thresholds count array.data_size() ELEMENTS (as the reference does),
|
||||
not allocation bytes; counters persist across encoder boundaries. The three
|
||||
new checks include dependent untracked GPU copies, two-queue event transfers
|
||||
and safe synthetic error-completion tests. The 48-test suite passes in both
|
||||
legacy retention modes. Only these new encoder tests use the independent
|
||||
queue. Complete operator access metadata, stream/evaluator integration and
|
||||
production routing are still open; this is not a model-performance receipt.
|
||||
|
||||
Operator scopes now connect the existing normalization and full MoE chain to
|
||||
the independent Concurrent encoder and its pooled untracked allocator. Explicit
|
||||
binding roles/data_size spans cover routing, sorting, gather/unsort, Gate/Up,
|
||||
SwiGLU, casts/norms, affine and gathered quantized projections, Split-K and both
|
||||
stock/fused experts plus the shared expert. Sort and split-reduction scratch is
|
||||
registered as backend temporaries. RMSNorm's default one and GatherSort's divisor
|
||||
are real scalar array bindings, not setBytes replacements. No shaders changed.
|
||||
The nine existing fixture groups (540 cases) execute through BOTH encoders with
|
||||
unchanged reference-output checks. This does not multiply independent fixtures.
|
||||
QSA routing, general array/donation semantics, evaluator/stream integration,
|
||||
production instrumentation and complete model/performance acceptance remain open.
|
||||
|
||||
GDN routing now uses the same typed dispatch/allocator path, including conv,
|
||||
mask/cache, Q/K normalization, compute_g/beta, recurrence, fused step and output.
|
||||
The direct native fused-step bypass is removed. Scalar operands and custom T
|
||||
are original scalar array inputs. Concatenate uses the original concurrent
|
||||
disjoint slice writes and dependency join; checked-input copies retain the
|
||||
original order and backend-temporary registration. No shader bodies changed.
|
||||
Four further existing groups (476 GDN cases) execute in both encoders, bringing
|
||||
the dual-encoder total to 1,016 existing cases. This remains operator-level
|
||||
correctness coverage, not whole-model scheduling or production parity.
|
||||
|
||||
QSA static/dynamic copies now bind explicit array data_size metadata rather
|
||||
than treating the copied region as the whole bound array. Dynamic offset arrays
|
||||
are inputs/backend temporaries. Zero fill, COW/General copies, compiled frontier
|
||||
operations and fused query/pool preparation carry access roles. KV concatenate
|
||||
uses the reference concurrent slice-write region. The existing 12 backing/copy
|
||||
and 26 compiled cache-window cases now run in both encoders (1,054 existing
|
||||
dual-encoder cases in total). Eager preparation, score/select operators and full
|
||||
graph/production integration remain open. Shader sources/geometries unchanged.
|
||||
|
||||
Eager QSA preparation now carries access roles and exact slice spans through
|
||||
RoPE, mean/RMS pooling and projections; all three RoPE concatenates use the
|
||||
reference concurrent writer regions. Scalar arithmetic uses array inputs.
|
||||
The fused selector and chunk concatenation also use the same encoder; its three
|
||||
frontiers are distinct scalar arrays and its private score plane remains a
|
||||
custom-kernel output, not a backend temporary. The existing 210 preparation,
|
||||
88 ongoing cache and 75 selector/chunk cases now execute in both encoders,
|
||||
bringing the existing dual-encoder total to 1,427 cases. Empty-array placeholder
|
||||
storage is unchanged; general zero-storage array semantics and the remaining
|
||||
score/eager-selection, evaluator and production paths are still open.
|
||||
|
||||
Both QSA prefill score producers and the connected score/top-k/chunk route now
|
||||
use the independent encoder. FP32 GEMV/Steel/NAX/Split-K, ReLU and reduction
|
||||
carry explicit array spans; Maximum/Divide scalars are array inputs. Matmul's
|
||||
output is allocated before transpose-copy preparation; those copies and the
|
||||
Split-K plane are registered as backend temporaries after their consumers.
|
||||
The MPP producer retains its original stride-aware inputs. Prefill top-k
|
||||
allocates outputs before checked-input copies, registers those copies after
|
||||
dispatch and uses three separately allocated frontier arrays. Shader bodies,
|
||||
specializations and dispatch geometries are unchanged. The existing 384 FP32
|
||||
score, 146 MPP/top-k/prefill and 96 mixed-producer cases now run in both encoders,
|
||||
bringing dual-encoder coverage to 2,053 existing cases. These manual fixture
|
||||
batches do not establish general evaluator/stream boundaries or production
|
||||
parity; eager selection/output and full graph integration remain open.
|
||||
|
||||
Eager QSA score masking, ranking and all output branches now carry explicit
|
||||
input/output array spans through the independent encoder. Scalar operands are
|
||||
real arrays; per-tile mx.eval(top_t) synchronizes that queue before constructing
|
||||
the next tile. Tiled index concatenation, decode-tail concatenation, rows-gather
|
||||
outputs and dense padding use the original concurrent slice-write regions.
|
||||
Dense padding is materialized before concatenation, and rows-gather builds its
|
||||
two separate Arange expressions; these restore previously collapsed operations.
|
||||
No Metal body or specialization changed.
|
||||
|
||||
The fixed-signature indexer graph now defers input/sibling Data on the active
|
||||
queue rather than retaining fallback descriptors in its returned state. The
|
||||
selected output's Data remains excluded, preserving donation. Compiled chunk
|
||||
concatenation joins concurrent writers and moving frontiers are four separate
|
||||
int32 arrays. Existing empty/nonempty Data-ownership checks run in both queues.
|
||||
The existing 540 eager cases and 200 ongoing indexer calls (including compiled
|
||||
routes) also execute in both encoders: 2,793 existing dual-encoder cases total.
|
||||
Explicit tile eval counts, output/state hashes and graph engagement are checked.
|
||||
These bounded fixture batches still do not apply the complete model-wide
|
||||
evaluator/stream scheduler and its automatic primitive-level commit policy.
|
||||
Production routing/instrumentation and whole-model performance acceptance remain
|
||||
open, as does general zero-storage array handling.
|
||||
|
||||
Primitive submission and task scheduler
|
||||
--------------------------------------
|
||||
|
||||
The fixed indexer tape now checks the original encoder thresholds after a whole
|
||||
primitive, never after an individual kernel dispatch. Input/sibling Data and
|
||||
backend holds are attached before a possible commit. Counted GPU tasks complete
|
||||
from the command callback; failed submissions balance ownership/accounting.
|
||||
The single-stream tape also applies the original active-task/memory pressure
|
||||
condition, finalizes its stream and waits for progress without inserting sleeps.
|
||||
|
||||
The Rust scheduler uses stdlib FIFO workers per CPU stream, earliest-error
|
||||
preservation, non-consuming cross-stream event error propagation and draining
|
||||
shutdown. CPU event waits/signals use the existing Metal shared-event bridge.
|
||||
Explicit tile/final fixture synchronizations are checked separately from
|
||||
automatic commits, with independent expectations for the two full-prefill
|
||||
indexer cases that cross the pinned Max data-size threshold.
|
||||
|
||||
CPU dispatch now counts every tenth operation, with completion as a separate
|
||||
FIFO task so failed work still completes its activity accounting. The CPU
|
||||
primitive cleanup task participates in that count and retains complete backend
|
||||
temporary descriptors until earlier work has run. CPU and GPU now share the
|
||||
input/sibling Data-selection function, excluding donated primary output Data.
|
||||
The private CPU temporary wrapper is Send only for drop-only worker ownership;
|
||||
Buffer itself remains !Send/!Sync. Five scheduler tests and the existing GPU
|
||||
ownership/connected-indexer checks pass, including unretained indexer execution.
|
||||
|
||||
An explicit runtime-owned stream registry now connects CPU/GPU encoders for the
|
||||
installed single-CPU/single-Metal backend. Defaults and template resolution are
|
||||
per thread/device. Local encoders are destroyed at thread exit; global streams
|
||||
allow sequential cross-thread use. Explicit clear preserves the reference's
|
||||
metadata/stale default handles and global CPU/GPU cleanup distinction.
|
||||
Registry locks do not cover encoding or waits. GPU selection reuses the same
|
||||
Rc/TLS encoder and allocator; exclusive ownership is checked before returning
|
||||
an encoder to storage. CPU/GPU events and finalize-all pressure handling use
|
||||
the existing scheduler and native bridge. Two new checks include actual
|
||||
two-queue copies, eleven blocked CPU tasks, thread cleanup and error unwinding.
|
||||
The 200 ongoing indexer cases now use a registered GPU stream in the independent
|
||||
path; the full model-free collection passes 57/57 normal and unretained.
|
||||
|
||||
Graph events now preserve per-copy values/origin streams while sharing native
|
||||
events and errors. Inter-stream fences use either the reference SharedEvent
|
||||
path or its opt-in Metal3/macOS15 fast path. Fast synchronization remains off
|
||||
by default. The existing pinned metallib supplies input_coherent, fence_update
|
||||
and fence_wait unchanged. Rust preserves array output registration, raw timestamp
|
||||
bindings, explicit update barriers, cross-device coherence and CPU SeqCst
|
||||
timestamp operations. Shared fence counts and per-dispatch snapshots are
|
||||
separate; completion/task ownership retains timestamp storage until work ends.
|
||||
The bridge now accepts empty dispatch grids, including the original zero-work
|
||||
coherence dispatch. Three new checks cover events/errors, partial-word/empty
|
||||
coherence and CPU/GPU/GPU/CPU transfer in both modes with a test-only deadlock
|
||||
rescue. All 60 model-free tests pass normal and unretained; no rescue fired.
|
||||
No shader body, metallib, inference default or production routing was changed.
|
||||
|
||||
This is still test-bound. Full array/evaluator cross-stream dependency construction, actual
|
||||
CPU model primitives, compile-cache cleanup integration and the complete model
|
||||
graph remain open. Production Qwen and performance acceptance remain unchanged;
|
||||
stream/operator checks do not establish whole-model parity.
|
||||
|
||||
Canonical dispatch now transfers deduplicated Rust root-allocation ownership
|
||||
to its existing Metal completion callback. This replaces the native resource
|
||||
set in that path and prevents physical release without adding Data aliases
|
||||
that would disable donation. Typed bindings also retain their descriptor/scalar
|
||||
borrow until dispatch. Early failures keep ownership with Rust; registration
|
||||
transfers it even when a submitted standalone command later reports failure.
|
||||
No shader, flush, commit, wait or additional completion-handler change.
|
||||
Normal and unretained collections pass 45/45 tests including native allocation,
|
||||
heap exhaustion, limits, residency and in-flight cache exclusion. This is not model parity.
|
||||
Binding input/output roles, barrier epochs, concurrent contexts, inter-encoder
|
||||
fences and reference commit thresholds still require the full encoder port.
|
||||
|
||||
Runtime-generated indexing shader
|
||||
--------------------------------
|
||||
|
||||
The same pinned runtime JIT-compiles gather_front instead of using the
|
||||
precompiled library. tools/mtplx-kernel-source.py resolves its two headers
|
||||
into metal/mtplx_qwen.metal without editing shader code and supplies the
|
||||
BF16/U32/FP32 template instantiations. Original notices and MIT attribution
|
||||
are retained. Header identities:
|
||||
|
||||
indexing/indexing.h:
|
||||
e820b8ee2b5132a97122780c12433ebb5100d8078d31e211d0429400a11415bb
|
||||
indexing/gather_front.h:
|
||||
64aacebf6576dfcd389383564fa1214bc87f2a091dd33cc64c598c5367ecab96
|
||||
|
||||
The generator requires the pinned runtime source checkout alongside the
|
||||
MTPLX reference checkout, named mtplx-runtime-0.32.2. The product requires
|
||||
only the generated shader resource, not this reference source checkout.
|
||||
|
||||
Runtime-generated SiLU shader
|
||||
----------------------------
|
||||
|
||||
tools/mtplx-jit-reference.py observes original runtime compilation without
|
||||
changing it. tests/fixtures/mtplx-silu-jit.json pins the generated19variants;
|
||||
full observed source SHA256:
|
||||
76cafb45db55a91efba66503dde59b37628ea360f220f35dd860f4ac3c3d0111
|
||||
|
||||
The shader exporter retains the generated computation and its original
|
||||
BF16 math, Sigmoid, Multiply, cast and stride helpers. Only host-name aliases
|
||||
change. All supporting header identities are enforced by the exporter and
|
||||
the emitted unit hashes by the focused Rust test. Copyright Apple Inc.; MIT
|
||||
as reproduced in MLX-LM-LICENSE.txt. This does not link a host runtime.
|
||||
|
||||
Runtime-generated GatherAxis shader
|
||||
-----------------------------------
|
||||
|
||||
The pinned GatherAxis source and generic elem_to_loc helper are retained
|
||||
unchanged, with BF16/U32 index instantiations for contiguous/strided inputs
|
||||
and int/int64 offsets. The router retains the strided last-ten-column view;
|
||||
no replacement top-k kernel is used. Header identities:
|
||||
|
||||
indexing/gather_axis.h:
|
||||
e1a745391ff4990f3f1ad75c5687c3b102dcdc4833d8fbbac38e10f54af29af4
|
||||
utils.h:
|
||||
5e1568e9edde9d05dbf86f68fa0d6c6240f2c32b973c7c6a76166b9c0d91543d
|
||||
|
||||
Copyright Apple Inc.; MIT as reproduced in MLX-LM-LICENSE.txt. Softmax,
|
||||
reduction, binary operations and index copies use the pinned metallib.
|
||||
|
||||
Runtime-generated SwiGLU shader
|
||||
------------------------------
|
||||
|
||||
The original compiled activations.swiglu used by non-sanitize-fused SwitchGLU
|
||||
and Qwen3NextMLP is captured separately from nn.silu. Full observed source SHA:
|
||||
bf78eee5cf96ea7c112c4e61546c12bcacf57fe512e0572182604cb94137510b.
|
||||
tests/fixtures/mtplx-swiglu-jit.json preserves all 19 original generated variants.
|
||||
The exporter changes only host aliases and reuses the already pinned BF16,
|
||||
Sigmoid, Multiply and cast/stride dependencies. Copyright Apple Inc.; MIT.
|
||||
|
||||
Runtime-generated compute_g shader and staged GDN
|
||||
------------------------------------------------
|
||||
|
||||
The original compiled gated_delta.compute_g is captured with the same observer:
|
||||
tests/fixtures/mtplx-compute_g-jit.json. Full observed source SHA256:
|
||||
34143a98046f8af5538767734fc169a5cab22a4920c26f9ba7ea45b8097152de.
|
||||
Receipt SHA256:
|
||||
701e2f54b7cb8bf97616f83256657e6c5e8cc8b46f4b65c559ccc030ab111dbf.
|
||||
All 19 variants retain their original BF16 Add/LogAddExp intermediates and
|
||||
FP32 final exponential. Only host aliases change. The exporter pins the
|
||||
additional Exp, Negative, Add, LogAddExp, Limits and log1p shader dependencies.
|
||||
|
||||
complex.h SHA256:
|
||||
16e8a815b2cbdb6070e0824e64fe33fccb6e918f1b84ea5c792bd89d33e57bf1.
|
||||
cexpf.h SHA256:
|
||||
88b6e15a52a5800d98d9bc6da840ca5cf70bf572fda136409580c1f17b1e0aab.
|
||||
The complex overload dependencies are retained unchanged, not used to add a
|
||||
complex-valued Qwen path. complex.h is Apple MIT. cexpf.h is Apache-2.0,
|
||||
Copyright Apple 2025, NVIDIA 2008-2013 and Filipe RNC Maia 2013. Its full original
|
||||
copyright/license notice remains embedded in the generated shader; the Apache
|
||||
license text is included in MTPLX-LICENSE.txt.
|
||||
|
||||
Stock depthwise Conv1D, copies, casts, reductions and elementary operations
|
||||
use the unchanged runtime metallib. Cache valid-length GatherAxis additionally
|
||||
instantiates the original signed INT32-index template; router indices remain
|
||||
UINT32. No shader body is replaced by a hand-written equivalent.
|
||||
|
||||
Original QSA indexer preparation
|
||||
-------------------------------
|
||||
|
||||
qsa_indexer_prepare.py SHA256:
|
||||
a77f6ca5ae805e729519c4629ae88b455a6dbf473a457a6e1c8219174eb59091.
|
||||
The exporter reads _prepare_queries_kernel and _pool_keys_kernel as AST data;
|
||||
it does not execute the model or kernel module. Both original source strings
|
||||
are unchanged. Header substitutions match the installed geometry: four query
|
||||
heads, width128, rotary64, ratio4, epsilon1e-6, attention scaling1. Includes are
|
||||
resolved at translation-unit scope; separate namespaces avoid collisions among
|
||||
the original header constants. Only entry-point declarations, host aliases and
|
||||
template instantiations are adapted. Stride metadata retains the original
|
||||
constant int64_t address space. Original Metal math and BF16 rounding remain.
|
||||
Copyright MTPLX; Apache-2.0, see MTPLX-LICENSE.txt and MTPLX-NOTICE.txt.
|
||||
|
||||
This is the preparation portion, not the full QSA indexer, selection,
|
||||
attention graph or production integration. Runtime frequencies are input buffers,
|
||||
not host replacements for the model's frequency construction.
|
||||
|
||||
QSACache/KVCache host lifecycle now uses the original scalar/vector/general
|
||||
copy and BF16/FP32 cast entries from this runtime, including positional writes,
|
||||
growth, strided restored state and the derived mirror. Rust distinguishes array
|
||||
object identity (__setitem__ overwrites its descriptor) from shared slice storage.
|
||||
Retained state aliases are checked against actual MTPLX cache operations, not
|
||||
assumed immutable. No additional shader bodies or runtime host library are used.
|
||||
The connected canonical cache remains test-only until product graph integration.
|
||||
|
||||
Original dynamic QSA selector
|
||||
-----------------------------
|
||||
|
||||
qsa_indexer_select.py SHA256:
|
||||
a3c74af27a7045c12f2893a8b7a91724c00d8a4148315c3165f3480c83016cf3.
|
||||
metal/mtplx-qsa-select.json preserves the original header, common body and all
|
||||
three epilogues (blocks, dense_mask, row_tokens), extracted without importing
|
||||
the model. Rust substitutes the original literal header parameters and supplies
|
||||
only the entry-point ABI. Tests additionally compare full generated header/body
|
||||
hashes against the actual MTPLX factory. H4/D128/ratio4 match the installed model;
|
||||
BF16/FP32 operands, backing capacity, top-k and TF32 remain specializations.
|
||||
Native compilation follows runtime 0.32.2 CustomKernel defaults: Safe math and
|
||||
its platform-selected Metal language version. No runtime host library is linked.
|
||||
The original 32MiB score-scratch chunk planner and typed output concatenation
|
||||
are connected to the cache/preparation port. General submission/concurrency,
|
||||
the complete eager indexer and production integration remain open.
|
||||
Copyright MTPLX; Apache-2.0, see MTPLX-LICENSE.txt and MTPLX-NOTICE.txt.
|
||||
|
||||
Original vectorized QSA prefill
|
||||
------------------------------
|
||||
|
||||
qsa_indexer_prefill.py SHA256:
|
||||
4d6fd428243c001746f69f8aed45991356772c2bd4a45586eb3c6813c91998d3.
|
||||
The same JSON export retains _MPP_SCORE_HEADER/_MPP_SCORE_SOURCE, the original
|
||||
top-k body and literal f-string header segments. Rust resolves only their named
|
||||
constants and provides entry-point ABI/type aliases. TensorOps tile layout,
|
||||
ordered per-head ReLU reduction, adaptive radix/insertion and all epilogues are
|
||||
unchanged. The original required General FP32 copy is used for non-contiguous
|
||||
score views; MPP input views keep their strides without added copies.
|
||||
The 128MiB producer-aware planner and score -> top-k -> concat chain are connected
|
||||
for the installed BF16/M5 geometry, including a 2K continuation from live cache.
|
||||
Full indexer branch routing, compiled graph bank and production integration
|
||||
remain open. Copyright MTPLX; Apache-2.0.
|
||||
|
||||
The general FP32 score expression now shares that prefill entry point. Rust
|
||||
ports runtime matmul.cpp's H4/D128 M5 Max routing: AsType Vector/General layout,
|
||||
check_transpose and broadcast copies, batch collapse, GEMV, regular Steel/NAX
|
||||
and both Split-K variants, original per-head Maximum, row/column Sum and Divide.
|
||||
All shader entries come from the unchanged pinned metallib; no shader body or
|
||||
host runtime library was added. The pooled cast is retained once across chunks;
|
||||
producer selection and the H4+1 workspace budget follow the reference.
|
||||
384 score cases and 96 connected selection cases are exact with real runtime
|
||||
MLX_ENABLE_TF32=0/1 in separate processes. These remain canonical correctness
|
||||
fixtures, not production integration or performance-parity evidence.
|
||||
|
||||
Original eager QSA selection
|
||||
---------------------------
|
||||
|
||||
The untiled QSAIndexer._select_eager score/top-k path uses the original
|
||||
runtime Arange, Add, integer Divide, Less, casts, Select, Subtract and
|
||||
ArgPartition (implemented by the pinned runtime as argsort). The chronological
|
||||
flash_prefill block epilogue adds original int32 Sort, int64 index conversion,
|
||||
bool GatherAxis and Select. The exporter adds only the required
|
||||
gather_axis<bool,int64_t,int,true,true> instantiation of the already preserved
|
||||
GatherAxis body; no body is changed. Its original file/unit hashes are unchanged.
|
||||
All untiled output epilogues are connected: dense mask (original bool
|
||||
ScatterAxis, repeat/concatenate and causal/tail mask), rows-gather (argsort-order
|
||||
tokens and validity), decode flash (chronological blocks, host tail bound) and
|
||||
decode gather (chronological tokens and variable-length tail). The flash branch
|
||||
retains precedence; neither decode branch evaluates the dead selected-mask DAG.
|
||||
Shared original cast/sort/vector dispatch helpers do not alter the shader bodies.
|
||||
|
||||
ScatterAxis adds these unchanged pinned Apple MIT runtime source units:
|
||||
- atomic.h, full-file SHA256:
|
||||
4c35ea2798a2335502865247aee878149fc9ada0d7e84c05d771baef0c7fcc60
|
||||
- reduction/ops.h None operation, full-file SHA256:
|
||||
78d06730fc9564a73944e7f1fe3897d25c8789b28a939bf418e1968db311da41
|
||||
- indexing/scatter_axis.h, full-file SHA256:
|
||||
43eabd0216101f8e32f5cdd19ce40b7f954564be27fad98a5e0fe345e7b94ce5
|
||||
Only include/pragma-once placement, namespace and the two required
|
||||
scatter_axis<bool,int64_t,int,None,false/true,true> instantiations are added
|
||||
outside the preserved bodies. The exporter guards full-file and body hashes.
|
||||
No host C/C++ runtime implementation is linked.
|
||||
|
||||
The 408 eager receipts tap actual QSAIndexer calls: 60 score/top-k, 24 prefill
|
||||
blocks, 144 dense masks, 108 rows-gather and 72 decode outputs. They include 2K
|
||||
queries, 65,536 blocks, tails 0/1/3 and separate real TF32-on/off processes.
|
||||
The tiled path shares the same score/rank functions, pooled FP32 input and tie
|
||||
vector. Each original mx.eval(top_t) is a synchronous command completion before
|
||||
the next tile, not an asynchronous flush. Only evaluated index views/backings
|
||||
are retained through the original GeneralGeneral uint32 concatenate. Index
|
||||
stride changes N -> K, independently of the N-strided validity. No new shader
|
||||
source or specialization is needed. Rust shares the original output-branch
|
||||
priority, including the tiled rows-gather exclusion and decode flash precedence.
|
||||
132 further actual-call receipts (108 tiled, 24 tile-off boundaries) bring the
|
||||
eager total to 540. They include observed reference eval row counts, 2K queries,
|
||||
65,536 blocks, tail 0/1, partial tiles and stride-2 FP32 query views.
|
||||
The existing GPU busy counters additionally verify one completed command buffer
|
||||
per observed reference tile eval, plus the final output batch.
|
||||
Full indexer routing, the general scheduling tape/allocator and production
|
||||
integration are not established by these checks.
|
||||
|
||||
Original eager QSA preparation
|
||||
------------------------------
|
||||
|
||||
The installed BF16 H4/D128 query and H1/D128 pool paths now use the original
|
||||
eager preparation expression as well as the fused custom-kernel branch.
|
||||
RMSNorm uses the pinned runtime kernel and its required General-Copy for sliced
|
||||
projection inputs. Pool mean is FP32 sum multiplied by 0.25 and cast to BF16
|
||||
before weighted RMSNorm. RoPE preserves all Arange, casts, concatenations,
|
||||
Cos/Sin, BF16 Negative, FP32 Multiply/Add and final BF16/pass-through stages.
|
||||
The installed rotary64/ratio4/eps1e-6/scaling1 contract is unchanged.
|
||||
All entries come from the unchanged pinned metallib; no shader body was added.
|
||||
|
||||
The existing projection and cache-extension entry points select either branch
|
||||
and retain the eager intermediates through their consuming operations. 105
|
||||
additional actual QSAIndexer receipts cover bare and quantized-projection
|
||||
preparation, including 2K rows, padded/stride-2 inputs and high positions. Another
|
||||
44 real QSACache/KVCache transitions cover eager pooling, capacity growth,
|
||||
reservation, trim, state aliases/restore and FP32 mirror rebuilding. These are
|
||||
not a complete indexer, production-integration or performance-parity receipt.
|
||||
|
||||
Connected non-compiled indexer entry
|
||||
----------------------------------
|
||||
|
||||
The Rust entry after compiled-route rejection now connects projection/supplied
|
||||
QK views, query preparation, raw/pool state and the original large-prefill,
|
||||
legacy-fused and eager selection order. All existing original shader dispatches
|
||||
are reused. Query preparation is dead when dense==sparse; KV.offset advances
|
||||
only in the subsequent Attention step. Shared return variants preserve the
|
||||
model-visible outputs while retaining other encoded kernel outputs.
|
||||
92 actual MTPLX indexer calls in 24 ongoing sequences verify lane choice, 82
|
||||
selection hashes, 268 raw/pool/mirror hashes, capacities/frontiers and command
|
||||
completion counts. They cover 2K rows, 32K history, supplied 704-stride QK views,
|
||||
prefill crossover, gather/flash priority and tiling. No compiled path is silently
|
||||
replaced. Compiled eligibility/core, attention, graph scheduling and production
|
||||
integration are still open; these are correctness, not performance receipts.
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,33 @@
|
||||
{
|
||||
"header_sha256": "2665a76463f3f6ee283c6a50b66e4a527318a114080b31441dfa900042097a39",
|
||||
"clamp": {
|
||||
"zero": {
|
||||
"max_new": 512,
|
||||
"max_start": 0,
|
||||
"source_sha256": "fd56a2d3bd76649775e853a28c41bd2a290255c7efbe4df471c6ed5e557b474e",
|
||||
"kernels": "[[host_name(\"Ei4IBroadcastBFi4ISubtractAEGi4IBroadcastCHi4IMaximumFGIi4IMinimumHDJi4OMinimumIG_VCCV_i4i4_13771019418134704434_contiguous\")]]\n[[kernel]] void Ei4IBroadcastBFi4ISubtractAEGi4IBroadcastCHi4IMaximumFGIi4IMinimumHDJi4OMinimumIG_VCCV_i4i4_13771019418134704434_contiguous(\n device const int32_t* A [[buffer(0)]],\n device const int32_t* B [[buffer(1)]],\n device int32_t* C [[buffer(2)]],\n constant const uint& size [[buffer(3)]],\n uint3 pos [[thread_position_in_grid]],\n uint3 grid [[threads_per_grid]]) {\n constexpr int N_ = 1;\n uint index = N_ * pos.x;\n int32_t tmp_A = A[index];\n auto tmp_D = static_cast<int32_t>(512);\n auto tmp_E = static_cast<int32_t>(0);\n int32_t tmp_B = B[index];\n int32_t tmp_F = cast_to<int32_t>(tmp_D);\n int32_t tmp_G = Subtract()(tmp_A, tmp_F);\n int32_t tmp_H = cast_to<int32_t>(tmp_E);\n int32_t tmp_I = Maximum()(tmp_G, tmp_H);\n int32_t tmp_J = Minimum()(tmp_I, tmp_B);\n int32_t tmp_C = Minimum()(tmp_J, tmp_H);\n C[index] = tmp_C;\n}\n"
|
||||
},
|
||||
"equal": {
|
||||
"max_new": 512,
|
||||
"max_start": 512,
|
||||
"source_sha256": "8f21de0cf23479618a4b0729740e197e50fec09f4f87fcb0caa8b8569ba8173a",
|
||||
"kernels": "[[host_name(\"Ei4IBroadcastBFi4ISubtractAEGi4IBroadcastCHi4IMaximumFGIi4IMinimumHDJi4OMinimumIE_VCCV_i4i4_13771019418134704434_contiguous\")]]\n[[kernel]] void Ei4IBroadcastBFi4ISubtractAEGi4IBroadcastCHi4IMaximumFGIi4IMinimumHDJi4OMinimumIE_VCCV_i4i4_13771019418134704434_contiguous(\n device const int32_t* A [[buffer(0)]],\n device const int32_t* B [[buffer(1)]],\n device int32_t* C [[buffer(2)]],\n constant const uint& size [[buffer(3)]],\n uint3 pos [[thread_position_in_grid]],\n uint3 grid [[threads_per_grid]]) {\n constexpr int N_ = 1;\n uint index = N_ * pos.x;\n int32_t tmp_A = A[index];\n auto tmp_D = static_cast<int32_t>(512);\n auto tmp_E = static_cast<int32_t>(0);\n int32_t tmp_B = B[index];\n int32_t tmp_F = cast_to<int32_t>(tmp_D);\n int32_t tmp_G = Subtract()(tmp_A, tmp_F);\n int32_t tmp_H = cast_to<int32_t>(tmp_E);\n int32_t tmp_I = Maximum()(tmp_G, tmp_H);\n int32_t tmp_J = Minimum()(tmp_I, tmp_B);\n int32_t tmp_C = Minimum()(tmp_J, tmp_F);\n C[index] = tmp_C;\n}\n"
|
||||
},
|
||||
"distinct": {
|
||||
"max_new": 1,
|
||||
"max_start": 1023,
|
||||
"source_sha256": "15dd362325e0d9ea44a8646c29ed70c7f0ed53b1dcc9fe44b4f65fa4a11823e7",
|
||||
"kernels": "[[host_name(\"Fi4IBroadcastBGi4ISubtractAFHi4IBroadcastCIi4IMaximumGHJi4IMinimumIDKi4IBroadcastELi4OMinimumJK_VCCVC_i4i4_1937821606537560661_contiguous\")]]\n[[kernel]] void Fi4IBroadcastBGi4ISubtractAFHi4IBroadcastCIi4IMaximumGHJi4IMinimumIDKi4IBroadcastELi4OMinimumJK_VCCVC_i4i4_1937821606537560661_contiguous(\n device const int32_t* A [[buffer(0)]],\n device const int32_t* B [[buffer(1)]],\n device int32_t* C [[buffer(2)]],\n constant const uint& size [[buffer(3)]],\n uint3 pos [[thread_position_in_grid]],\n uint3 grid [[threads_per_grid]]) {\n constexpr int N_ = 1;\n uint index = N_ * pos.x;\n int32_t tmp_A = A[index];\n auto tmp_D = static_cast<int32_t>(1);\n auto tmp_E = static_cast<int32_t>(0);\n int32_t tmp_B = B[index];\n auto tmp_F = static_cast<int32_t>(1023);\n int32_t tmp_G = cast_to<int32_t>(tmp_D);\n int32_t tmp_H = Subtract()(tmp_A, tmp_G);\n int32_t tmp_I = cast_to<int32_t>(tmp_E);\n int32_t tmp_J = Maximum()(tmp_H, tmp_I);\n int32_t tmp_K = Minimum()(tmp_J, tmp_B);\n int32_t tmp_L = cast_to<int32_t>(tmp_F);\n int32_t tmp_C = Minimum()(tmp_K, tmp_L);\n C[index] = tmp_C;\n}\n"
|
||||
}
|
||||
},
|
||||
"multiply": {
|
||||
"source_sha256": "1cf792edbbd886d156b68f5082563d7335278c98ad8c996d547b19fd7cf125b0",
|
||||
"kernels": "[[host_name(\"Ci4IBroadcastBDi4OMultiplyAC_VC_i4_2169371982377735806_contiguous\")]]\n[[kernel]] void Ci4IBroadcastBDi4OMultiplyAC_VC_i4_2169371982377735806_contiguous(\n device const int32_t* A [[buffer(0)]],\n device int32_t* B [[buffer(1)]],\n constant const uint& size [[buffer(2)]],\n uint3 pos [[thread_position_in_grid]],\n uint3 grid [[threads_per_grid]]) {\n constexpr int N_ = 1;\n uint index = N_ * pos.x;\n int32_t tmp_A = A[index];\n auto tmp_C = static_cast<int32_t>(4);\n int32_t tmp_D = cast_to<int32_t>(tmp_C);\n int32_t tmp_B = Multiply()(tmp_A, tmp_D);\n B[index] = tmp_B;\n}\n"
|
||||
},
|
||||
"check_distinct": {
|
||||
"max_new": 257,
|
||||
"max_start": 255,
|
||||
"source_sha256": "c952e6026e18bc5c5cc6f3b78852607d9efa8ce7a9b083d439515c95c1620ef9",
|
||||
"kernels": "[[host_name(\"Fi4IBroadcastBGi4ISubtractAFHi4IBroadcastCIi4IMaximumGHJi4IMinimumIDKi4IBroadcastELi4OMinimumJK_VCCVC_i4i4_7252438397961030063_contiguous\")]]\n[[kernel]] void Fi4IBroadcastBGi4ISubtractAFHi4IBroadcastCIi4IMaximumGHJi4IMinimumIDKi4IBroadcastELi4OMinimumJK_VCCVC_i4i4_7252438397961030063_contiguous(\n device const int32_t* A [[buffer(0)]],\n device const int32_t* B [[buffer(1)]],\n device int32_t* C [[buffer(2)]],\n constant const uint& size [[buffer(3)]],\n uint3 pos [[thread_position_in_grid]],\n uint3 grid [[threads_per_grid]]) {\n constexpr int N_ = 1;\n uint index = N_ * pos.x;\n int32_t tmp_A = A[index];\n auto tmp_D = static_cast<int32_t>(257);\n auto tmp_E = static_cast<int32_t>(0);\n int32_t tmp_B = B[index];\n auto tmp_F = static_cast<int32_t>(255);\n int32_t tmp_G = cast_to<int32_t>(tmp_D);\n int32_t tmp_H = Subtract()(tmp_A, tmp_G);\n int32_t tmp_I = cast_to<int32_t>(tmp_E);\n int32_t tmp_J = Maximum()(tmp_H, tmp_I);\n int32_t tmp_K = Minimum()(tmp_J, tmp_B);\n int32_t tmp_L = cast_to<int32_t>(tmp_F);\n int32_t tmp_C = Minimum()(tmp_K, tmp_L);\n C[index] = tmp_C;\n}\n"
|
||||
}
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
Binary file not shown.
+10419
File diff suppressed because it is too large
Load Diff
+570
-40
@@ -340,8 +340,8 @@ kernel void kernel_qwen_affine_qmv_fast(
|
||||
|
||||
// MLX affine_qmv_wide for the verifier's S=2..4 matrices. M5-class GPUs
|
||||
// route these shapes here instead of running one independent QMV per row.
|
||||
template <ushort bits, ushort group_size>
|
||||
static inline void qwen_affine_qmv_wide_impl(
|
||||
template <ushort bits, ushort group_size, ushort rows>
|
||||
static inline void qwen_affine_qmv_wide_rows_impl(
|
||||
constant qwen_kernel_args &args,
|
||||
device float *out,
|
||||
device const float *x,
|
||||
@@ -361,12 +361,11 @@ static inline void qwen_affine_qmv_wide_impl(
|
||||
simd_group * outputs_per_simdgroup + simd_row;
|
||||
const uint in_dim = args.u[0];
|
||||
const uint out_dim = args.u[1];
|
||||
const uint rows = args.u[4];
|
||||
const uint row = min(output_row, out_dim - 1u);
|
||||
const uint groups_per_row = in_dim / group_size;
|
||||
const ulong packed_row = (ulong)row * in_dim * bits / 8u;
|
||||
const ulong parameter_row = (ulong)row * groups_per_row;
|
||||
float result[4] = {0.0f, 0.0f, 0.0f, 0.0f};
|
||||
float result[rows] = {0.0f};
|
||||
|
||||
for (uint quant_group = k_lane; quant_group < groups_per_row;
|
||||
quant_group += k_lanes) {
|
||||
@@ -374,11 +373,13 @@ static inline void qwen_affine_qmv_wide_impl(
|
||||
scales, args.u[14], parameter_row + quant_group));
|
||||
const float bias = qwen_bf16(qwen_weight_u16(
|
||||
biases, args.u[15], parameter_row + quant_group));
|
||||
#pragma unroll
|
||||
for (uint chunk = 0u; chunk < group_size / sub; chunk++) {
|
||||
const uint column = quant_group * group_size + chunk * sub;
|
||||
const ulong weight_byte = (ulong)args.u[13] + packed_row +
|
||||
(ulong)column * bits / 8u;
|
||||
float weights[sub];
|
||||
#pragma unroll
|
||||
for (uint index = 0u; index < sub; index++) {
|
||||
uint quantized;
|
||||
if constexpr (bits == 4) {
|
||||
@@ -389,8 +390,10 @@ static inline void qwen_affine_qmv_wide_impl(
|
||||
}
|
||||
weights[index] = scale * (float)quantized + bias;
|
||||
}
|
||||
#pragma unroll
|
||||
for (uint vector = 0u; vector < rows; vector++) {
|
||||
float sum = 0.0f;
|
||||
#pragma unroll
|
||||
for (uint index = 0u; index < sub; index++) {
|
||||
sum += x[(ulong)vector * in_dim + column + index] * weights[index];
|
||||
}
|
||||
@@ -411,6 +414,35 @@ static inline void qwen_affine_qmv_wide_impl(
|
||||
}
|
||||
}
|
||||
|
||||
// MTPLX specializes vecs_per_tg: keeping the accumulator indices constant
|
||||
// avoids a runtime-indexed register array in the verifier's inner loop.
|
||||
template <ushort bits, ushort group_size>
|
||||
static inline void qwen_affine_qmv_wide_impl(
|
||||
constant qwen_kernel_args &args,
|
||||
device float *out,
|
||||
device const float *x,
|
||||
device const uchar *packed,
|
||||
device const uchar *scales,
|
||||
device const uchar *biases,
|
||||
uint group,
|
||||
uint simd_group,
|
||||
uint lane) {
|
||||
switch (args.u[4]) {
|
||||
case 2:
|
||||
qwen_affine_qmv_wide_rows_impl<bits, group_size, 2>(
|
||||
args, out, x, packed, scales, biases, group, simd_group, lane);
|
||||
break;
|
||||
case 3:
|
||||
qwen_affine_qmv_wide_rows_impl<bits, group_size, 3>(
|
||||
args, out, x, packed, scales, biases, group, simd_group, lane);
|
||||
break;
|
||||
case 4:
|
||||
qwen_affine_qmv_wide_rows_impl<bits, group_size, 4>(
|
||||
args, out, x, packed, scales, biases, group, simd_group, lane);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
template <ushort bits, ushort group_size>
|
||||
kernel void kernel_qwen_affine_qmv_batch_fast(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
@@ -583,23 +615,111 @@ kernel qwen_affine_pair_qmv_wide_b8g64 kernel_qwen_affine_pair_qmv_wide<8, 64>;
|
||||
#ifdef DS4_METAL_HAS_TENSOR
|
||||
using namespace mpp::tensor_ops;
|
||||
|
||||
template <ushort bits, ushort group_size, ushort tile_m>
|
||||
kernel void kernel_qwen_affine_qmm_mpp(
|
||||
kernel void kernel_qwen_route_map(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device uint *counts [[buffer(1)]],
|
||||
device const uint *selected [[buffer(2)]],
|
||||
device uint *route_map [[buffer(3)]],
|
||||
device uchar *work [[buffer(4)]],
|
||||
threadgroup ushort *staged [[threadgroup(0)]],
|
||||
uint expert [[thread_index_in_threadgroup]],
|
||||
uint threads [[threads_per_threadgroup]]) {
|
||||
uint count = 0u;
|
||||
device uint *expert_map = route_map + (ulong)expert * args.u[4];
|
||||
for (uint first_row = 0u; first_row < args.u[4]; first_row += threads) {
|
||||
const uint row = first_row + expert;
|
||||
if (row < args.u[4]) {
|
||||
for (uint slot = 0u; slot < args.u[8]; slot++) {
|
||||
staged[expert * args.u[8] + slot] =
|
||||
(ushort)selected[(ulong)row * args.u[8] + slot];
|
||||
}
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
const uint valid = min(threads, args.u[4] - first_row);
|
||||
for (uint local_row = 0u; local_row < valid; local_row++) {
|
||||
for (uint slot = 0u; slot < args.u[8]; slot++) {
|
||||
if (staged[local_row * args.u[8] + slot] == expert) {
|
||||
expert_map[count++] = (first_row + local_row) * args.u[8] + slot;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
}
|
||||
counts[expert] = count;
|
||||
threadgroup_barrier(mem_flags::mem_device);
|
||||
uint route_base = 0u;
|
||||
for (uint other = 0u; other < expert; other++) route_base += counts[other];
|
||||
device uint *route_offsets = (device uint *)(work + 8);
|
||||
device uint *inverse = route_offsets + 512;
|
||||
route_offsets[expert] = route_base;
|
||||
for (uint index = 0u; index < count; index++) {
|
||||
inverse[expert_map[index]] = route_base + index;
|
||||
}
|
||||
staged[expert] = (ushort)((count + 63u) / 64u);
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
uint work_base = 0u;
|
||||
for (uint other = 0u; other < expert; other++) {
|
||||
work_base += staged[other];
|
||||
}
|
||||
device uint *work_count = (device uint *)work;
|
||||
const ulong items_offset =
|
||||
(8ul + 512ul * 4ul + (ulong)args.u[4] * args.u[8] * 4ul + 7ul) & ~7ul;
|
||||
device uint2 *work_items = (device uint2 *)(work + items_offset);
|
||||
for (uint tile = 0u; tile < staged[expert]; tile++) {
|
||||
work_items[work_base + tile] = uint2(expert, tile * 64u);
|
||||
}
|
||||
if (expert + 1u == threads) {
|
||||
work_count[0] = work_base + staged[expert];
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_route_gather_rows(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device bfloat *out [[buffer(1)]],
|
||||
device const float *x [[buffer(2)]],
|
||||
device const uchar *work [[buffer(3)]],
|
||||
uint2 gid [[thread_position_in_grid]]) {
|
||||
const uint column = gid.x;
|
||||
const uint route = gid.y;
|
||||
if (column >= args.u[0] || route >= args.u[4] * args.u[8]) return;
|
||||
device const uint *inverse = (device const uint *)(work + 8) + 512;
|
||||
const uint sorted = inverse[route];
|
||||
out[(ulong)sorted * args.u[0] + column] =
|
||||
(bfloat)x[(ulong)(route / args.u[8]) * args.u[0] + column];
|
||||
}
|
||||
|
||||
template <ushort bits, ushort group_size>
|
||||
kernel void kernel_qwen_affine_gather_qmm_mpp(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device const float *x [[buffer(2)]],
|
||||
device const uint *route_map [[buffer(3)]],
|
||||
device const uchar *work [[buffer(4)]],
|
||||
device const uchar *packed [[buffer(5)]],
|
||||
device const uchar *scales [[buffer(6)]],
|
||||
device const uchar *biases [[buffer(7)]],
|
||||
device const uint *counts [[buffer(8)]],
|
||||
uint2 group [[threadgroup_position_in_grid]],
|
||||
uint tid [[thread_index_in_threadgroup]]) {
|
||||
constexpr uint tile_m = 64u;
|
||||
constexpr uint tile_n = 64u;
|
||||
constexpr uint tile_k = 64u;
|
||||
constexpr uint threads = 128u;
|
||||
device const uint *work_count = (device const uint *)work;
|
||||
if (group.y >= work_count[0]) return;
|
||||
const ulong items_offset =
|
||||
(8ul + 512ul * 4ul + (ulong)args.u[4] * args.u[8] * 4ul + 7ul) & ~7ul;
|
||||
device const uint2 *work_items = (device const uint2 *)(work + items_offset);
|
||||
const uint2 item = work_items[group.y];
|
||||
const uint expert = item.x;
|
||||
const uint first_m = item.y;
|
||||
if (first_m >= counts[expert]) return;
|
||||
const uint first_n = group.x * tile_n;
|
||||
device const uint *expert_map = route_map + (ulong)expert * args.u[4];
|
||||
threadgroup bfloat xs[tile_m * tile_k];
|
||||
threadgroup bfloat ws[tile_n * tile_k];
|
||||
const uint first_m = group.y * tile_m;
|
||||
const uint first_n = group.x * tile_n;
|
||||
constexpr auto descriptor = matmul2d_descriptor(
|
||||
tile_m, tile_n, tile_k, false, true, false,
|
||||
matmul2d_descriptor::mode::multiply_accumulate);
|
||||
@@ -609,28 +729,29 @@ kernel void kernel_qwen_affine_qmm_mpp(
|
||||
tensor<threadgroup bfloat, dextents<int32_t, 2>, tensor_inline>,
|
||||
float>();
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) {
|
||||
accum[i] = 0.0f;
|
||||
}
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) accum[i] = 0.0f;
|
||||
|
||||
for (uint first_k = 0u; first_k < args.u[0]; first_k += tile_k) {
|
||||
for (uint index = tid; index < tile_m * tile_k; index += threads) {
|
||||
const uint row = index / tile_k;
|
||||
const uint column = index % tile_k;
|
||||
const uint input_m = first_m + row;
|
||||
xs[index] = input_m < args.u[4]
|
||||
? (bfloat)x[(ulong)input_m * args.u[0] + first_k + column]
|
||||
const bool valid = first_m + row < counts[expert];
|
||||
const uint route = valid ? expert_map[first_m + row] : 0u;
|
||||
const uint input_row = args.u[5] != 0u ? route : route / args.u[8];
|
||||
xs[index] = valid
|
||||
? (bfloat)x[(ulong)input_row * args.u[0] + first_k + column]
|
||||
: bfloat(0.0f);
|
||||
}
|
||||
for (uint index = tid; index < tile_n * tile_k; index += threads) {
|
||||
const uint row = index / tile_k;
|
||||
const uint column = index % tile_k;
|
||||
const uint weight_n = first_n + row;
|
||||
const ulong table_row = (ulong)expert * args.u[1] + weight_n;
|
||||
ws[index] = weight_n < args.u[1]
|
||||
? (bfloat)qwen_quant_weight(
|
||||
packed, scales, biases,
|
||||
args.u[13], args.u[14], args.u[15],
|
||||
weight_n, first_k + column,
|
||||
table_row, first_k + column,
|
||||
args.u[0], bits, group_size)
|
||||
: bfloat(0.0f);
|
||||
}
|
||||
@@ -648,15 +769,267 @@ kernel void kernel_qwen_affine_qmm_mpp(
|
||||
const auto index = accum.get_multidimensional_index(i);
|
||||
const uint output_n = first_n + index[0];
|
||||
const uint output_m = first_m + index[1];
|
||||
if (output_n < args.u[1] && output_m < args.u[4]) {
|
||||
const float value = accum[i];
|
||||
out[(ulong)output_m * args.u[1] + output_n] = args.u[11] != 0u
|
||||
? qwen_round_bf16(value)
|
||||
: value;
|
||||
if (output_n < args.u[1] && output_m < counts[expert]) {
|
||||
const uint route = expert_map[output_m];
|
||||
out[(ulong)route * args.u[1] + output_n] =
|
||||
qwen_round_bf16(accum[i]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <ushort bits, ushort group_size>
|
||||
kernel void kernel_qwen_affine_sorted_qmm_mpp(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device bfloat *x [[buffer(2)]],
|
||||
device const uint *route_map [[buffer(3)]],
|
||||
device const uchar *work [[buffer(4)]],
|
||||
device const uchar *packed [[buffer(5)]],
|
||||
device const uchar *scales [[buffer(6)]],
|
||||
device const uchar *biases [[buffer(7)]],
|
||||
device const uint *counts [[buffer(8)]],
|
||||
uint2 group [[threadgroup_position_in_grid]],
|
||||
uint tid [[thread_index_in_threadgroup]]) {
|
||||
constexpr uint tile_m = 64u;
|
||||
constexpr uint tile_n = 64u;
|
||||
constexpr uint tile_k = 32u;
|
||||
constexpr uint threads = 128u;
|
||||
device const uint *work_count = (device const uint *)work;
|
||||
if (group.y >= work_count[0]) return;
|
||||
const ulong items_offset =
|
||||
(8ul + 512ul * 4ul + (ulong)args.u[4] * args.u[8] * 4ul + 7ul) & ~7ul;
|
||||
device const uint2 *work_items = (device const uint2 *)(work + items_offset);
|
||||
const uint2 item = work_items[group.y];
|
||||
const uint expert = item.x;
|
||||
const uint local_m = item.y;
|
||||
if (local_m >= counts[expert]) return;
|
||||
device const uint *route_offsets = (device const uint *)(work + 8);
|
||||
const uint first_m = route_offsets[expert] + local_m;
|
||||
const uint first_n = group.x * tile_n;
|
||||
threadgroup bfloat ws[2u * tile_n * tile_k];
|
||||
auto weights0 = tensor<threadgroup bfloat, dextents<int32_t, 2>, tensor_inline>(
|
||||
ws, dextents<int32_t, 2>(tile_k, tile_n));
|
||||
auto weights1 = tensor<threadgroup bfloat, dextents<int32_t, 2>, tensor_inline>(
|
||||
ws + tile_n * tile_k, dextents<int32_t, 2>(tile_k, tile_n));
|
||||
auto activations = tensor<device bfloat, dextents<int32_t, 2>, tensor_inline>(
|
||||
x, dextents<int32_t, 2>(args.u[0], args.u[6]),
|
||||
array<int, 2>({1, (int)args.u[0]}));
|
||||
constexpr auto descriptor = matmul2d_descriptor(
|
||||
tile_m, tile_n, tile_k, false, true, true,
|
||||
matmul2d_descriptor::mode::multiply_accumulate);
|
||||
matmul2d<descriptor, execution_simdgroups<4>> multiply;
|
||||
auto accum = multiply.template get_destination_cooperative_tensor<
|
||||
decltype(activations), decltype(weights0), float>();
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) accum[i] = 0.0f;
|
||||
|
||||
auto stage_weights = [&](uint first_k, threadgroup bfloat *target) {
|
||||
constexpr uint per_word = 32u / bits;
|
||||
constexpr uint words_per_row = tile_k / per_word;
|
||||
for (uint index = tid; index < tile_n * words_per_row; index += threads) {
|
||||
const uint row = index / words_per_row;
|
||||
const uint column = (index % words_per_row) * per_word;
|
||||
const uint weight_n = first_n + row;
|
||||
const ulong table_row = (ulong)expert * args.u[1] + weight_n;
|
||||
uint word = 0u;
|
||||
float scale = 0.0f, bias = 0.0f;
|
||||
if (weight_n < args.u[1]) {
|
||||
word = qwen_weight_u32(packed, args.u[13],
|
||||
table_row * (args.u[0] / per_word) + (first_k + column) / per_word);
|
||||
const ulong quant_group = table_row * (args.u[0] / group_size)
|
||||
+ (first_k + column) / group_size;
|
||||
scale = qwen_bf16(qwen_weight_u16(scales, args.u[14], quant_group));
|
||||
bias = qwen_bf16(qwen_weight_u16(biases, args.u[15], quant_group));
|
||||
}
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint element = 0u; element < per_word; element++) {
|
||||
const uint quant = (word >> (element * bits)) & ((1u << bits) - 1u);
|
||||
target[row * tile_k + column + element] =
|
||||
(bfloat)fma((float)quant, scale, bias);
|
||||
}
|
||||
}
|
||||
};
|
||||
stage_weights(0u, ws);
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
uint selected_weights = 0u;
|
||||
for (uint first_k = 0u; first_k < args.u[0]; first_k += tile_k) {
|
||||
auto weight_tile = selected_weights ? weights1 : weights0;
|
||||
auto activation_tile = activations.slice(first_k, first_m);
|
||||
multiply.run(activation_tile, weight_tile, accum);
|
||||
const uint next_k = first_k + tile_k;
|
||||
if (next_k < args.u[0]) {
|
||||
selected_weights ^= 1u;
|
||||
stage_weights(next_k, selected_weights ? ws + tile_n * tile_k : ws);
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
}
|
||||
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) {
|
||||
const auto index = accum.get_multidimensional_index(i);
|
||||
const uint output_n = first_n + index[0];
|
||||
const uint local_output_m = local_m + index[1];
|
||||
if (output_n < args.u[1] && local_output_m < counts[expert]) {
|
||||
out[(ulong)(first_m + index[1]) * args.u[1] + output_n] =
|
||||
qwen_round_bf16(accum[i]);
|
||||
}
|
||||
}
|
||||
(void)route_map;
|
||||
}
|
||||
|
||||
typedef decltype(kernel_qwen_affine_sorted_qmm_mpp<4, 32>) qwen_affine_sorted_qmm_mpp_b4g32;
|
||||
typedef decltype(kernel_qwen_affine_sorted_qmm_mpp<4, 64>) qwen_affine_sorted_qmm_mpp_b4g64;
|
||||
typedef decltype(kernel_qwen_affine_sorted_qmm_mpp<8, 64>) qwen_affine_sorted_qmm_mpp_b8g64;
|
||||
template [[host_name("kernel_qwen_affine_sorted_qmm_mpp_b4g32")]]
|
||||
kernel qwen_affine_sorted_qmm_mpp_b4g32 kernel_qwen_affine_sorted_qmm_mpp<4, 32>;
|
||||
template [[host_name("kernel_qwen_affine_sorted_qmm_mpp_b4g64")]]
|
||||
kernel qwen_affine_sorted_qmm_mpp_b4g64 kernel_qwen_affine_sorted_qmm_mpp<4, 64>;
|
||||
template [[host_name("kernel_qwen_affine_sorted_qmm_mpp_b8g64")]]
|
||||
kernel qwen_affine_sorted_qmm_mpp_b8g64 kernel_qwen_affine_sorted_qmm_mpp<8, 64>;
|
||||
|
||||
typedef decltype(kernel_qwen_affine_gather_qmm_mpp<4, 32>) qwen_affine_gather_qmm_mpp_b4g32;
|
||||
typedef decltype(kernel_qwen_affine_gather_qmm_mpp<4, 64>) qwen_affine_gather_qmm_mpp_b4g64;
|
||||
typedef decltype(kernel_qwen_affine_gather_qmm_mpp<8, 64>) qwen_affine_gather_qmm_mpp_b8g64;
|
||||
|
||||
template [[host_name("kernel_qwen_affine_gather_qmm_mpp_b4g32")]]
|
||||
kernel qwen_affine_gather_qmm_mpp_b4g32 kernel_qwen_affine_gather_qmm_mpp<4, 32>;
|
||||
template [[host_name("kernel_qwen_affine_gather_qmm_mpp_b4g64")]]
|
||||
kernel qwen_affine_gather_qmm_mpp_b4g64 kernel_qwen_affine_gather_qmm_mpp<4, 64>;
|
||||
template [[host_name("kernel_qwen_affine_gather_qmm_mpp_b8g64")]]
|
||||
kernel qwen_affine_gather_qmm_mpp_b8g64 kernel_qwen_affine_gather_qmm_mpp<8, 64>;
|
||||
|
||||
template <ushort bits, ushort group_size, ushort tile_m>
|
||||
kernel void kernel_qwen_affine_qmm_mpp(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device float *x [[buffer(2)]],
|
||||
device const uchar *packed [[buffer(5)]],
|
||||
device const uchar *scales [[buffer(6)]],
|
||||
device const uchar *biases [[buffer(7)]],
|
||||
uint2 group [[threadgroup_position_in_grid]],
|
||||
uint tid [[thread_index_in_threadgroup]]) {
|
||||
constexpr uint tile_n = 64u;
|
||||
constexpr uint tile_k = 32u;
|
||||
constexpr uint threads = 128u;
|
||||
threadgroup bfloat ws[2u * tile_n * tile_k];
|
||||
const uint m_tiles = (args.u[4] + tile_m - 1u) / tile_m;
|
||||
const uint partition = group.y / m_tiles;
|
||||
const uint first_m = (group.y % m_tiles) * tile_m;
|
||||
const uint first_n = group.x * tile_n;
|
||||
const uint partition_k = args.u[0] / max(args.u[8], 1u);
|
||||
const uint start_k = partition * partition_k;
|
||||
const uint end_k = start_k + partition_k;
|
||||
out += partition * args.u[4] * args.u[1];
|
||||
auto weights0 = tensor<threadgroup bfloat, dextents<int32_t, 2>, tensor_inline>(
|
||||
ws, dextents<int32_t, 2>(tile_k, tile_n));
|
||||
auto weights1 = tensor<threadgroup bfloat, dextents<int32_t, 2>, tensor_inline>(
|
||||
ws + tile_n * tile_k, dextents<int32_t, 2>(tile_k, tile_n));
|
||||
auto activations = tensor<device float, dextents<int32_t, 2>, tensor_inline>(
|
||||
x, dextents<int32_t, 2>(args.u[0], args.u[4]),
|
||||
array<int, 2>({1, (int)args.u[0]}));
|
||||
constexpr auto descriptor = matmul2d_descriptor(
|
||||
tile_m, tile_n, tile_k, false, true, true,
|
||||
matmul2d_descriptor::mode::multiply_accumulate);
|
||||
matmul2d<descriptor, execution_simdgroups<4>> multiply;
|
||||
auto accum = multiply.template get_destination_cooperative_tensor<
|
||||
decltype(activations), decltype(weights0),
|
||||
float>();
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) {
|
||||
accum[i] = 0.0f;
|
||||
}
|
||||
|
||||
auto stage_weights = [&](uint first_k, threadgroup bfloat *target) {
|
||||
// Decode a packed word once, reusing its scale/bias across its values,
|
||||
// as in the sorted QMM loader and MTPLX's quantized matrix loader.
|
||||
constexpr uint per_word = 32u / bits;
|
||||
constexpr uint words_per_row = tile_k / per_word;
|
||||
for (uint index = tid; index < tile_n * words_per_row; index += threads) {
|
||||
const uint row = index / words_per_row;
|
||||
const uint column = (index % words_per_row) * per_word;
|
||||
const uint weight_n = first_n + row;
|
||||
uint word = 0u;
|
||||
float scale = 0.0f, bias = 0.0f;
|
||||
if (weight_n < args.u[1]) {
|
||||
word = qwen_weight_u32(packed, args.u[13],
|
||||
(ulong)weight_n * (args.u[0] / per_word) + (first_k + column) / per_word);
|
||||
const ulong quant_group = (ulong)weight_n * (args.u[0] / group_size)
|
||||
+ (first_k + column) / group_size;
|
||||
scale = qwen_bf16(qwen_weight_u16(scales, args.u[14], quant_group));
|
||||
bias = qwen_bf16(qwen_weight_u16(biases, args.u[15], quant_group));
|
||||
}
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint element = 0u; element < per_word; element++) {
|
||||
const uint quant = (word >> (element * bits)) & ((1u << bits) - 1u);
|
||||
target[row * tile_k + column + element] =
|
||||
(bfloat)fma((float)quant, scale, bias);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
stage_weights(start_k, ws);
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
uint selected_weights = 0u;
|
||||
for (uint first_k = start_k; first_k < end_k; first_k += tile_k) {
|
||||
auto weight_tile = selected_weights ? weights1 : weights0;
|
||||
auto activation_tile = activations.slice(first_k, first_m);
|
||||
multiply.run(activation_tile, weight_tile, accum);
|
||||
const uint next_k = first_k + tile_k;
|
||||
if (next_k < end_k) {
|
||||
selected_weights ^= 1u;
|
||||
stage_weights(
|
||||
next_k,
|
||||
selected_weights ? ws + tile_n * tile_k : ws);
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
}
|
||||
|
||||
auto output = tensor<device float, dextents<int32_t, 2>, tensor_inline>(
|
||||
out, dextents<int32_t, 2>(args.u[1], args.u[4]),
|
||||
array<int, 2>({1, (int)args.u[1]}));
|
||||
if (args.u[11] != 0u) {
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0u; i < accum.get_capacity(); i++) {
|
||||
accum[i] = qwen_round_bf16(accum[i]);
|
||||
}
|
||||
}
|
||||
auto output_tile = output.slice(first_n, first_m);
|
||||
accum.store(output_tile);
|
||||
}
|
||||
|
||||
// Match MTPLX's BF16 column-reduction order, including intermediate rounding.
|
||||
// Each simdgroup handles one output; small reductions use eight strided lanes,
|
||||
// larger reductions use the reference's 32-lane BF16 simd reduction.
|
||||
kernel void kernel_qwen_affine_splitk_reduce(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device const float *parts [[buffer(2)]],
|
||||
uint group [[threadgroup_position_in_grid]],
|
||||
uint simd_group [[simdgroup_index_in_threadgroup]],
|
||||
uint lane [[thread_index_in_simdgroup]]) {
|
||||
const uint index = group * 2u + simd_group;
|
||||
const uint stride = args.u[4] * args.u[1];
|
||||
if (index >= stride) return;
|
||||
const uint count = args.u[8];
|
||||
const uint lanes = count < 32u ? min(count, 8u) : 32u;
|
||||
bfloat total = bfloat(0.0f);
|
||||
for (uint p = lane; p < count; p += lanes) {
|
||||
if (lane < lanes) total = bfloat(float(total) + parts[p * stride + index]);
|
||||
}
|
||||
if (count < 32u) {
|
||||
bfloat sum = total;
|
||||
for (uint p = 1u; p < lanes; p++) {
|
||||
const float next = simd_shuffle(float(total), p);
|
||||
sum = bfloat(float(sum) + float(next));
|
||||
}
|
||||
if (lane == 0u) out[index] = float(sum);
|
||||
} else {
|
||||
// MTPLX's bfloat16_t overload reduces in float, then rounds once.
|
||||
const bfloat sum = bfloat(simd_sum(float(total)));
|
||||
if (lane == 0u) out[index] = float(sum);
|
||||
}
|
||||
}
|
||||
|
||||
typedef decltype(kernel_qwen_affine_qmm_mpp<4, 32, 32>) qwen_affine_qmm_mpp_b4g32_bm32;
|
||||
typedef decltype(kernel_qwen_affine_qmm_mpp<4, 64, 32>) qwen_affine_qmm_mpp_b4g64_bm32;
|
||||
typedef decltype(kernel_qwen_affine_qmm_mpp<8, 64, 32>) qwen_affine_qmm_mpp_b8g64_bm32;
|
||||
@@ -1209,6 +1582,33 @@ kernel void kernel_qwen_weighted_sum10(
|
||||
out[index] = value;
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_weighted_sum10_sorted(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device const float *experts [[buffer(2)]],
|
||||
device const float *weights [[buffer(3)]],
|
||||
device const uchar *work [[buffer(4)]],
|
||||
uint index [[thread_position_in_grid]]) {
|
||||
if (index >= args.u[0]) return;
|
||||
const uint hidden = args.u[1];
|
||||
const uint row = index / hidden;
|
||||
const uint column = index % hidden;
|
||||
device const uint *inverse = (device const uint *)(work + 8) + 512;
|
||||
const ulong routes_base = (ulong)row * args.u[8];
|
||||
float weighted[10];
|
||||
for (uint slot = 0; slot < 10u; slot++) {
|
||||
const ulong route = routes_base + slot;
|
||||
weighted[slot] = qwen_round_bf16(
|
||||
experts[(ulong)inverse[route] * hidden + column] * weights[route]);
|
||||
}
|
||||
float value = qwen_round_bf16(weighted[0] + weighted[8]);
|
||||
value = qwen_round_bf16(value + qwen_round_bf16(weighted[1] + weighted[9]));
|
||||
for (uint slot = 2; slot < 8u; slot++) {
|
||||
value = qwen_round_bf16(value + weighted[slot]);
|
||||
}
|
||||
out[index] = value;
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_affine_embedding(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
@@ -1959,7 +2359,9 @@ kernel void kernel_qwen_gdn_conv_norm(
|
||||
: qwen_round_bf16(normalized);
|
||||
}
|
||||
|
||||
// MTPLX fused_gdn_conv_norm_rows, including its in-window convolution tail.
|
||||
// MTPLX fused_gdn_conv_norm_rows at S<=6. Larger prefills preserve the
|
||||
// Conv1d -> SiLU -> L2 BF16 boundaries and parallelize independent rows.
|
||||
// In the parallel path state_out must not alias the input state.
|
||||
kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *q_out [[buffer(1)]],
|
||||
@@ -1972,12 +2374,12 @@ kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
uint tid [[thread_index_in_threadgroup]],
|
||||
uint lane [[thread_index_in_simdgroup]],
|
||||
uint simd_group [[simdgroup_index_in_threadgroup]],
|
||||
uint group [[threadgroup_position_in_grid]]) {
|
||||
uint2 group [[threadgroup_position_in_grid]]) {
|
||||
constexpr uint width = 10240u;
|
||||
constexpr uint key_width = 2048u;
|
||||
constexpr uint dim = 128u;
|
||||
constexpr float inverse_scale = 0.08838834764831845f;
|
||||
const uint channel = group * 1024u + tid;
|
||||
const uint channel = group.x * 1024u + tid;
|
||||
if (channel >= width) return;
|
||||
|
||||
threadgroup float values[1024];
|
||||
@@ -1988,7 +2390,10 @@ kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
const float w3 = qwen_bf16(qwen_weight_u16(weight, args.u[14], (ulong)channel * 4u + 3u));
|
||||
const bool value_channel = channel >= 2u * key_width;
|
||||
|
||||
for (uint row = 0u; row < args.u[4]; row++) {
|
||||
const bool prefill = args.u[4] > 6u;
|
||||
const uint first_row = prefill ? group.y : 0u;
|
||||
const uint end_row = prefill ? first_row + 1u : args.u[4];
|
||||
for (uint row = first_row; row < end_row; row++) {
|
||||
const float x0 = row < 3u
|
||||
? qwen_bf16(state[(ulong)row * width + channel])
|
||||
: qkv[(ulong)(row - 3u) * width + channel];
|
||||
@@ -2000,7 +2405,10 @@ kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
: qkv[(ulong)(row - 1u) * width + channel];
|
||||
const float x3 = qkv[(ulong)row * width + channel];
|
||||
const float convolved = w0 * x0 + w1 * x1 + w2 * x2 + w3 * x3;
|
||||
const float activated = convolved / (1.0f + exp(-convolved));
|
||||
const float rounded = qwen_round_bf16(convolved);
|
||||
const float activated = prefill
|
||||
? qwen_round_bf16(rounded * qwen_silu_sigmoid_bf16(rounded))
|
||||
: convolved / (1.0f + exp(-convolved));
|
||||
if (value_channel) {
|
||||
v_out[(ulong)row * (width - 2u * key_width) + channel - 2u * key_width] =
|
||||
qwen_round_bf16(activated);
|
||||
@@ -2013,10 +2421,12 @@ kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
const uint first = (simd_group / 4u) * 4u;
|
||||
sum = partial[first] + partial[first + 1u] +
|
||||
partial[first + 2u] + partial[first + 3u];
|
||||
const float normalized = values[tid] * rsqrt(sum + 1.0e-6f);
|
||||
const float normalized = values[tid] * (prefill
|
||||
? precise::rsqrt(sum + 1.0e-6f) : rsqrt(sum + 1.0e-6f));
|
||||
if (channel < key_width) {
|
||||
q_out[(ulong)row * key_width + channel] =
|
||||
qwen_round_bf16(normalized * inverse_scale);
|
||||
prefill ? qwen_round_bf16(qwen_round_bf16(normalized) * 0.08837890625f)
|
||||
: qwen_round_bf16(normalized * inverse_scale);
|
||||
} else {
|
||||
k_out[(ulong)row * key_width + channel - key_width] =
|
||||
qwen_round_bf16(normalized);
|
||||
@@ -2024,6 +2434,7 @@ kernel void kernel_qwen_gdn_conv_norm_rows(
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
}
|
||||
|
||||
if (prefill && group.y != 0u) return;
|
||||
for (uint tail = 0u; tail < 3u; tail++) {
|
||||
const uint sequence = args.u[4] + tail;
|
||||
const float value = sequence < 3u
|
||||
@@ -2609,18 +3020,28 @@ kernel void kernel_qwen_gdn_norm_gate(
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_swiglu(
|
||||
template <typename T>
|
||||
kernel void kernel_qwen_swiglu_typed(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device T *out [[buffer(1)]],
|
||||
device const float *gate [[buffer(2)]],
|
||||
device const float *up [[buffer(3)]],
|
||||
uint index [[thread_position_in_grid]]) {
|
||||
if (index >= args.u[0]) return;
|
||||
const float activated = qwen_round_bf16(
|
||||
gate[index] * qwen_silu_sigmoid_bf16(gate[index]));
|
||||
out[index] = qwen_round_bf16(activated * up[index]);
|
||||
out[index] = (T)qwen_round_bf16(activated * up[index]);
|
||||
}
|
||||
|
||||
template [[host_name("kernel_qwen_swiglu")]]
|
||||
kernel void kernel_qwen_swiglu_typed<float>(
|
||||
constant qwen_kernel_args &, device float *, device const float *,
|
||||
device const float *, uint);
|
||||
template [[host_name("kernel_qwen_swiglu_bf16")]]
|
||||
kernel void kernel_qwen_swiglu_typed<bfloat>(
|
||||
constant qwen_kernel_args &, device bfloat *, device const float *,
|
||||
device const float *, uint);
|
||||
|
||||
kernel void kernel_qwen_unpack_gdn_inputs(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *qkv [[buffer(1)]],
|
||||
@@ -3017,10 +3438,18 @@ kernel void kernel_qwen_qsa_scores(
|
||||
device float *scores [[buffer(1)]],
|
||||
device const float *query [[buffer(2)]],
|
||||
device const ushort *pooled [[buffer(3)]],
|
||||
uint block [[thread_position_in_grid]]) {
|
||||
uint2 group [[thread_position_in_grid]]) {
|
||||
const uint block = group.x;
|
||||
const uint row = group.y;
|
||||
const uint dim = args.u[0];
|
||||
if (block >= args.u[2]) return;
|
||||
const ulong score_index = (ulong)row * args.u[2] + block;
|
||||
if (args.u[5] != 0u && block >= (args.u[3] + row) / args.u[5]) {
|
||||
scores[score_index] = -INFINITY;
|
||||
return;
|
||||
}
|
||||
float score = 0.0f;
|
||||
query += (ulong)row * args.u[1] * dim;
|
||||
for (uint head = 0; head < args.u[1]; head++) {
|
||||
float head_score = 0.0f;
|
||||
for (uint i = 0; i < dim; i++) {
|
||||
@@ -3030,14 +3459,14 @@ kernel void kernel_qwen_qsa_scores(
|
||||
}
|
||||
score += max(head_score, 0.0f);
|
||||
}
|
||||
scores[block] = score * args.f[0];
|
||||
scores[score_index] = score * args.f[0];
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_qsa_sort_blocks(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device int *selected [[buffer(1)]],
|
||||
uint gid [[thread_position_in_grid]]) {
|
||||
if (gid != 0u) return;
|
||||
uint row [[thread_position_in_grid]]) {
|
||||
selected += (ulong)row * args.u[0];
|
||||
for (uint i = 1; i < args.u[0]; i++) {
|
||||
const int value = selected[i];
|
||||
uint j = i;
|
||||
@@ -3049,15 +3478,45 @@ kernel void kernel_qwen_qsa_sort_blocks(
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_qsa_mask(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device uchar *mask [[buffer(1)]],
|
||||
device const int *selected [[buffer(2)]],
|
||||
uint row [[threadgroup_position_in_grid]],
|
||||
uint lane [[thread_index_in_threadgroup]]) {
|
||||
if (row >= args.u[4]) return;
|
||||
const uint tokens = args.u[0];
|
||||
const uint visible = args.u[8] + row;
|
||||
const uint tail_start = (visible / args.u[5]) * args.u[5];
|
||||
const ulong mask_base = (ulong)row * tokens;
|
||||
for (uint token = lane; token < tokens; token += args.u[12]) {
|
||||
mask[mask_base + token] = token >= tail_start && token < visible;
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_device);
|
||||
if (lane < args.u[6]) {
|
||||
const int block = selected[(ulong)row * args.u[6] + lane];
|
||||
if (block >= 0) {
|
||||
const uint token0 = (uint)block * args.u[5];
|
||||
if (token0 < tail_start) {
|
||||
for (uint within = 0u; within < args.u[5]; within++) {
|
||||
mask[mask_base + token0 + within] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_sparse_attention(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
device const float *query [[buffer(2)]],
|
||||
device const ushort *cache [[buffer(3)]],
|
||||
device const int *selected [[buffer(4)]],
|
||||
uint head [[threadgroup_position_in_grid]],
|
||||
uint tid [[thread_position_in_threadgroup]],
|
||||
uint2 group [[threadgroup_position_in_grid]],
|
||||
uint tid [[thread_index_in_threadgroup]],
|
||||
uint lane [[thread_index_in_simdgroup]]) {
|
||||
const uint head = group.x;
|
||||
const uint row = group.y;
|
||||
const uint heads = args.u[0];
|
||||
const uint kv_heads = args.u[1];
|
||||
const uint dim = args.u[2];
|
||||
@@ -3071,7 +3530,14 @@ kernel void kernel_qwen_sparse_attention(
|
||||
threadgroup float sum_exp_scores[simdgroups];
|
||||
threadgroup float outputs[simdgroups * 256u];
|
||||
const float attention_scale = precise::rsqrt((float)dim);
|
||||
q[tid] = query[(ulong)head * dim + tid] * attention_scale;
|
||||
const ulong query_base = ((ulong)row * heads + head) * dim;
|
||||
const ulong output_base = query_base;
|
||||
selected += (ulong)row * args.u[4];
|
||||
const uint tokens = args.u[8] != 0u ? args.u[8] + row : args.u[3];
|
||||
const uint tail_start = args.u[8] != 0u
|
||||
? (tokens / args.u[5]) * args.u[5]
|
||||
: args.u[6];
|
||||
q[tid] = query[query_base + tid] * attention_scale;
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
float result[values_per_lane] = {0.0f};
|
||||
@@ -3101,8 +3567,8 @@ kernel void kernel_qwen_sparse_attention(
|
||||
max_score = new_max;
|
||||
}
|
||||
}
|
||||
for (uint token = args.u[6] + simd_group;
|
||||
token < args.u[3];
|
||||
for (uint token = tail_start + simd_group;
|
||||
token < tokens;
|
||||
token += simdgroups) {
|
||||
const ulong base = (ulong)token * kv_heads * dim * 2u +
|
||||
(ulong)kv_head * dim + column;
|
||||
@@ -3142,7 +3608,7 @@ kernel void kernel_qwen_sparse_attention(
|
||||
merged_sum += sum_exp_scores[group] * weight;
|
||||
merged_value += outputs[group * dim + tid] * weight;
|
||||
}
|
||||
out[(ulong)head * dim + tid] = qwen_round_bf16(merged_value / merged_sum);
|
||||
out[output_base + tid] = qwen_round_bf16(merged_value / merged_sum);
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_dense_attention_masked(
|
||||
@@ -3393,6 +3859,52 @@ kernel void kernel_qwen_attention_fallback_softmax(
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_attention_fallback_softmax_wide(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device ushort *out [[buffer(1)]],
|
||||
device const ushort *input [[buffer(2)]],
|
||||
uint group [[threadgroup_position_in_grid]],
|
||||
uint lid [[thread_position_in_threadgroup]],
|
||||
uint lane [[thread_index_in_simdgroup]],
|
||||
uint simd_group [[simdgroup_index_in_threadgroup]]) {
|
||||
threadgroup float local_max[32];
|
||||
threadgroup float local_normalizer[32];
|
||||
const uint tokens = args.u[3];
|
||||
const uint threads = args.u[12];
|
||||
const ulong base = (ulong)group * tokens;
|
||||
float maximum = -FLT_MAX;
|
||||
for (uint token = lid; token < tokens; token += threads) {
|
||||
maximum = max(maximum, qwen_bf16(input[base + token]));
|
||||
}
|
||||
maximum = simd_max(maximum);
|
||||
if (lane == 0u) local_max[simd_group] = maximum;
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
if (simd_group == 0u) {
|
||||
maximum = simd_max(local_max[lane]);
|
||||
if (lane == 0u) local_max[0] = maximum;
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
maximum = local_max[0];
|
||||
|
||||
float normalizer = 0.0f;
|
||||
for (uint token = lid; token < tokens; token += threads) {
|
||||
normalizer += fast::exp(qwen_bf16(input[base + token]) - maximum);
|
||||
}
|
||||
normalizer = simd_sum(normalizer);
|
||||
if (lane == 0u) local_normalizer[simd_group] = normalizer;
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
if (simd_group == 0u) {
|
||||
normalizer = simd_sum(local_normalizer[lane]);
|
||||
if (lane == 0u) local_normalizer[0] = normalizer;
|
||||
}
|
||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||
normalizer = 1.0f / local_normalizer[0];
|
||||
for (uint token = lid; token < tokens; token += threads) {
|
||||
out[base + token] = qwen_to_bf16(
|
||||
fast::exp(qwen_bf16(input[base + token]) - maximum) * normalizer);
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_attention_fallback_output(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device float *out [[buffer(1)]],
|
||||
@@ -3449,6 +3961,24 @@ kernel void kernel_qwen_prepare_dense_kv(
|
||||
values[target] = cache[source + width];
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_prepare_dense_kv_f16(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device half *keys [[buffer(1)]],
|
||||
device const ushort *cache [[buffer(2)]],
|
||||
device half *values [[buffer(3)]],
|
||||
uint2 gid [[thread_position_in_grid]]) {
|
||||
const uint width = args.u[0];
|
||||
const uint dim = args.u[1];
|
||||
const uint tokens = args.u[3];
|
||||
if (gid.x >= width || gid.y >= tokens) return;
|
||||
const uint head = gid.x / dim;
|
||||
const uint column = gid.x % dim;
|
||||
const ulong source = (ulong)gid.y * width * 2u + gid.x;
|
||||
const ulong target = ((ulong)head * tokens + gid.y) * dim + column;
|
||||
keys[target] = (half)qwen_bf16(cache[source]);
|
||||
values[target] = (half)qwen_bf16(cache[source + width]);
|
||||
}
|
||||
|
||||
kernel void kernel_qwen_float_to_bf16(
|
||||
constant qwen_kernel_args &args [[buffer(0)]],
|
||||
device ushort *out [[buffer(1)]],
|
||||
|
||||
Reference in New Issue
Block a user