mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Build the GGML IQ packing kernels as a single CUDA extension (#2462)
### What does this PR do? Type of change: Code refactoring `#2448` added the GGML IQ packing kernels as **two** torch extensions, `modelopt_cuda_ext_iq1_s` and `modelopt_cuda_ext_iq2_xs`. This merges them into one, `modelopt_cuda_ext_ggml`. The existing per-extension split in `extensions.py` exists for reasons that don't apply to the IQ formats: `get_cuda_ext` gates on CUDA `>=11` while `_fp8`/`_mx` gate on `>=11.8`, and `_mx` needs `--use_fast_math`, which must not reach the base `tensor_quant` kernels. `get_cuda_ext_iq1_s` and `get_cuda_ext_iq2_xs` differed in none of that — same `>=11.8` gate, same `-O3` flags, same `common.cuh` — so the split only compiled the shared header twice, ran nvcc twice, and grew the loader, `__getattr__`, and `precompile()` once per format. With IQ2_XXS / IQ3_S / IQ4_NL plausibly following, that scales badly. Changes: - New `ggml/ggml.cpp` holds both host-side validation wrappers and the single `PYBIND11_MODULE`, binding `iq1_s_pack` and `iq2_xs_pack` (previously each module exported a bare `pack`). Deletes `ggml/iq1_s.cpp` and `ggml/iq2_xs.cpp`; the validation logic and docstrings carry over unchanged. - `get_cuda_ext_iq1_s` + `get_cuda_ext_iq2_xs` → `get_cuda_ext_ggml`, which builds `ggml.cpp`, `iq1_s.cu`, and `iq2_xs.cu` together. The retry-on-`raise_if_failed` semantics of the old getters are preserved. - Each format keeps its kernels in its own translation unit, so adding a format is a new `.cu` plus one `module.def` — no new extension, loader, or `precompile()` line. No caller outside `extensions.py` and its tests referenced the old getters on `main`, so nothing else changes. **Note for the follow-up PRs in the `#2448` series (`#2446`/`#2447`/`#2449`): the codec layer should call `get_cuda_ext_ggml().iq1_s_pack(...)` / `.iq2_xs_pack(...)` instead of `get_cuda_ext_iq1_s().pack(...)` / `get_cuda_ext_iq2_xs().pack(...)`.** ### Usage ```python from modelopt.torch.quantization.extensions import get_cuda_ext_ggml ext = get_cuda_ext_ggml(raise_if_failed=True) iq1_s_payload = ext.iq1_s_pack(weight, iq1s_grid) # uint8 [numel / 256, 50] iq2_xs_payload = ext.iq2_xs_pack(weight, iq2xs_grid, scales) # uint8 [numel / 256, 74] ``` ### Testing Ran on a single H200 NVL (TRT-LLM `1.3.0rc27.dev202609170000` container), building the merged extension from scratch: - `pytest tests/gpu/_extensions/test_torch_extensions.py` — **24 passed** (6:44). This is the full existing IQ suite (zero-block layout, encode, dtype rejection, row-straddling rejection, invalid/negative-zero scales, byte-exact dtype equivalence, and the brute-force optimality round-trip) reparametrized onto the merged module, plus the untouched `modelopt_cuda_ext` / `_fp8` / `_mx` load tests. - Verified `precompile()` loads all four extensions and that the merged module exports exactly `iq1_s_pack` and `iq2_xs_pack` with the expected arities. - Off-GPU: compiled the three sources directly and linked them into one `.so` to confirm no duplicate-symbol collisions between the two `.cu` translation units. - `pre-commit run --files ...` passes on all changed files (ruff, mypy, clang-format, bandit, license headers). ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ — the removed getters were added in `#2448` (merged today, unreleased) and have no callers outside this file's own tests. - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: N/A — no new code or dependencies; the moved wrappers keep their original attribution. - Did you write any new necessary tests?: ✅ — existing coverage reparametrized onto the merged module; no behavior change to test. - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: N/A — internal refactor of an unreleased, not-yet-wired-up API. - Did you get Claude approval on this PR?: ❌ — not yet run. ### Additional Information Follow-up to #2448. Merge before the remaining PRs in that series (#2446, #2447, #2449) land, so the codec layer is written against `get_cuda_ext_ggml` and no rename is needed afterwards. 🤖 Generated with [Claude Code](https://claude.com/claude-code) <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **New Features** - Added IQ1_S packing support through the GGML CUDA extension. - Added a unified GGML extension loader for IQ1_S and IQ2_XS packing. - Improved extension loading reliability when a cached extension is unavailable. - **Changes** - Renamed the IQ2_XS packing binding from `pack` to `iq2_xs_pack`. - Consolidated IQ1_S and IQ2_XS extension access under the shared GGML loader. - Updated GPU validation and coverage to use the unified extension interface. <!-- end of auto-generated comment: release notes by coderabbit.ai --> Signed-off-by: Chenjie Luo <chenjiel@nvidia.com> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
542012d4d5
commit
b16356776b
+22
-1
@@ -15,10 +15,24 @@
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
// Every GGML IQ format shares common.cuh, the same CUDA version gate, and the same build flags,
|
||||
// so they compile into one extension and bind here. Each format keeps its kernels in its own
|
||||
// translation unit and exposes a single host entry point.
|
||||
|
||||
#include "common.cuh"
|
||||
|
||||
at::Tensor iq1_s_pack_cuda(at::Tensor input, at::Tensor grid);
|
||||
at::Tensor iq2_xs_pack_cuda(at::Tensor input, at::Tensor grid, at::Tensor scales);
|
||||
|
||||
namespace {
|
||||
|
||||
at::Tensor iq1_s_pack(at::Tensor input, at::Tensor grid) {
|
||||
TORCH_CHECK(input.is_cuda(), "IQ1_S packing requires a CUDA input");
|
||||
TORCH_CHECK(grid.is_cuda(), "IQ1_S packing requires a CUDA grid");
|
||||
modelopt::ggml::check_pack_inputs("IQ1_S", input, grid, modelopt::ggml::kIq1sEntries);
|
||||
return iq1_s_pack_cuda(input.contiguous(), grid.contiguous());
|
||||
}
|
||||
|
||||
at::Tensor iq2_xs_pack(at::Tensor input, at::Tensor grid, at::Tensor scales) {
|
||||
TORCH_CHECK(input.is_cuda(), "IQ2_XS packing requires a CUDA input");
|
||||
TORCH_CHECK(grid.is_cuda(), "IQ2_XS packing requires a CUDA grid");
|
||||
@@ -39,8 +53,15 @@ at::Tensor iq2_xs_pack(at::Tensor input, at::Tensor grid, at::Tensor scales) {
|
||||
return iq2_xs_pack_cuda(input.contiguous(), grid.contiguous(), scales.contiguous());
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
|
||||
module.def("pack", &iq2_xs_pack,
|
||||
module.def("iq1_s_pack", &iq1_s_pack,
|
||||
"Pack a non-empty float32, float64, float16, or bfloat16 CUDA tensor whose innermost "
|
||||
"dimension is a multiple of 256. The grid must be float32 [2048, 8]. Returns uint8 "
|
||||
"[numel / 256, 50] on the input device. Non-finite input elements are treated as "
|
||||
"zero during packing, and finite elements outside the float32 range saturate.");
|
||||
module.def("iq2_xs_pack", &iq2_xs_pack,
|
||||
"Pack a non-empty float32, float64, float16, or bfloat16 CUDA tensor whose innermost "
|
||||
"dimension is a multiple of 256. The grid must be float32 [512, 8] holding "
|
||||
"non-negative codebook magnitudes, and scales must be finite non-negative float16 "
|
||||
@@ -1,35 +0,0 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: Apache-2.0
|
||||
*
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#include "common.cuh"
|
||||
|
||||
at::Tensor iq1_s_pack_cuda(at::Tensor input, at::Tensor grid);
|
||||
|
||||
at::Tensor iq1_s_pack(at::Tensor input, at::Tensor grid) {
|
||||
TORCH_CHECK(input.is_cuda(), "IQ1_S packing requires a CUDA input");
|
||||
TORCH_CHECK(grid.is_cuda(), "IQ1_S packing requires a CUDA grid");
|
||||
modelopt::ggml::check_pack_inputs("IQ1_S", input, grid, modelopt::ggml::kIq1sEntries);
|
||||
return iq1_s_pack_cuda(input.contiguous(), grid.contiguous());
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
|
||||
module.def("pack", &iq1_s_pack,
|
||||
"Pack a non-empty float32, float64, float16, or bfloat16 CUDA tensor whose innermost "
|
||||
"dimension is a multiple of 256. The grid must be float32 [2048, 8]. Returns uint8 "
|
||||
"[numel / 256, 50] on the input device. Non-finite input elements are treated as "
|
||||
"zero during packing, and finite elements outside the float32 range saturate.");
|
||||
}
|
||||
@@ -22,8 +22,7 @@ from modelopt.torch.utils import load_cpp_extension
|
||||
__all__ = [
|
||||
"get_cuda_ext",
|
||||
"get_cuda_ext_fp8",
|
||||
"get_cuda_ext_iq1_s",
|
||||
"get_cuda_ext_iq2_xs",
|
||||
"get_cuda_ext_ggml",
|
||||
"get_cuda_ext_mx",
|
||||
"precompile",
|
||||
]
|
||||
@@ -80,36 +79,29 @@ def get_cuda_ext_mx(raise_if_failed: bool = False):
|
||||
return get_cuda_ext_mx.extension # type:ignore[attr-defined]
|
||||
|
||||
|
||||
def get_cuda_ext_iq1_s(raise_if_failed: bool = False):
|
||||
"""Return the GGML-compatible IQ1_S packing extension."""
|
||||
if not hasattr(get_cuda_ext_iq1_s, "extension") or (
|
||||
raise_if_failed and get_cuda_ext_iq1_s.extension is None
|
||||
def get_cuda_ext_ggml(raise_if_failed: bool = False):
|
||||
"""Return the GGML-compatible IQ packing extension, exposing one packer per IQ format.
|
||||
|
||||
The formats share their packing helpers, CUDA version requirement, and build flags, so they
|
||||
build as a single extension: ``iq1_s_pack(input, grid)`` and
|
||||
``iq2_xs_pack(input, grid, scales)``.
|
||||
"""
|
||||
if not hasattr(get_cuda_ext_ggml, "extension") or (
|
||||
raise_if_failed and get_cuda_ext_ggml.extension is None
|
||||
):
|
||||
get_cuda_ext_iq1_s.extension = load_cpp_extension( # type:ignore[attr-defined]
|
||||
name="modelopt_cuda_ext_iq1_s",
|
||||
sources=[kernels_ggml / "iq1_s.cpp", kernels_ggml / "iq1_s.cu"],
|
||||
get_cuda_ext_ggml.extension = load_cpp_extension( # type:ignore[attr-defined]
|
||||
name="modelopt_cuda_ext_ggml",
|
||||
sources=[
|
||||
kernels_ggml / "ggml.cpp",
|
||||
kernels_ggml / "iq1_s.cu",
|
||||
kernels_ggml / "iq2_xs.cu",
|
||||
],
|
||||
cuda_version_specifiers=">=11.8",
|
||||
fail_msg="IQ1_S CUDA packing extension is unavailable.",
|
||||
fail_msg="GGML IQ CUDA packing extension is unavailable.",
|
||||
extra_cuda_cflags=["-O3"],
|
||||
raise_if_failed=raise_if_failed,
|
||||
)
|
||||
return get_cuda_ext_iq1_s.extension # type:ignore[attr-defined]
|
||||
|
||||
|
||||
def get_cuda_ext_iq2_xs(raise_if_failed: bool = False):
|
||||
"""Return the GGML-compatible IQ2_XS packing extension."""
|
||||
if not hasattr(get_cuda_ext_iq2_xs, "extension") or (
|
||||
raise_if_failed and get_cuda_ext_iq2_xs.extension is None
|
||||
):
|
||||
get_cuda_ext_iq2_xs.extension = load_cpp_extension( # type:ignore[attr-defined]
|
||||
name="modelopt_cuda_ext_iq2_xs",
|
||||
sources=[kernels_ggml / "iq2_xs.cpp", kernels_ggml / "iq2_xs.cu"],
|
||||
cuda_version_specifiers=">=11.8",
|
||||
fail_msg="IQ2_XS CUDA packing extension is unavailable.",
|
||||
extra_cuda_cflags=["-O3"],
|
||||
raise_if_failed=raise_if_failed,
|
||||
)
|
||||
return get_cuda_ext_iq2_xs.extension # type:ignore[attr-defined]
|
||||
return get_cuda_ext_ggml.extension # type:ignore[attr-defined]
|
||||
|
||||
|
||||
def __getattr__(name):
|
||||
@@ -119,10 +111,8 @@ def __getattr__(name):
|
||||
return get_cuda_ext_fp8()
|
||||
elif name == "cuda_ext_mx":
|
||||
return get_cuda_ext_mx()
|
||||
elif name == "cuda_ext_iq1_s":
|
||||
return get_cuda_ext_iq1_s()
|
||||
elif name == "cuda_ext_iq2_xs":
|
||||
return get_cuda_ext_iq2_xs()
|
||||
elif name == "cuda_ext_ggml":
|
||||
return get_cuda_ext_ggml()
|
||||
else:
|
||||
raise AttributeError(f"module {__name__} has no attribute {name}")
|
||||
|
||||
@@ -132,5 +122,4 @@ def precompile():
|
||||
print(get_cuda_ext())
|
||||
print(get_cuda_ext_fp8())
|
||||
print(get_cuda_ext_mx())
|
||||
print(get_cuda_ext_iq1_s())
|
||||
print(get_cuda_ext_iq2_xs())
|
||||
print(get_cuda_ext_ggml())
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
# limitations under the License.
|
||||
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import NamedTuple
|
||||
|
||||
import pytest
|
||||
@@ -39,12 +38,8 @@ def test_cuda_ext_mx():
|
||||
assert ext.get_cuda_ext_mx() is not None
|
||||
|
||||
|
||||
def test_cuda_ext_iq1_s():
|
||||
assert ext.get_cuda_ext_iq1_s() is not None
|
||||
|
||||
|
||||
def test_cuda_ext_iq2_xs():
|
||||
assert ext.get_cuda_ext_iq2_xs() is not None
|
||||
def test_cuda_ext_ggml():
|
||||
assert ext.get_cuda_ext_ggml() is not None
|
||||
|
||||
|
||||
def _generator():
|
||||
@@ -53,9 +48,10 @@ def _generator():
|
||||
|
||||
|
||||
class _IqFormat(NamedTuple):
|
||||
"""One GGML IQ packing extension and the format constants its contract is defined by."""
|
||||
"""One GGML IQ packer and the format constants its contract is defined by."""
|
||||
|
||||
get_extension: Callable
|
||||
# Name the packer is bound under on the shared GGML extension.
|
||||
packer: str
|
||||
entries: int
|
||||
payload_bytes: int
|
||||
needs_scales: bool
|
||||
@@ -68,11 +64,9 @@ class _IqFormat(NamedTuple):
|
||||
|
||||
|
||||
_IQ_EXTENSIONS = (
|
||||
pytest.param(_IqFormat("iq1_s_pack", 2048, 50, False, (-1.0, 0.0, 1.0), 16.875), id="iq1_s"),
|
||||
pytest.param(
|
||||
_IqFormat(ext.get_cuda_ext_iq1_s, 2048, 50, False, (-1.0, 0.0, 1.0), 16.875), id="iq1_s"
|
||||
),
|
||||
pytest.param(
|
||||
_IqFormat(ext.get_cuda_ext_iq2_xs, 512, 74, True, (8.0, 25.0, 43.0), 166.625),
|
||||
_IqFormat("iq2_xs_pack", 512, 74, True, (8.0, 25.0, 43.0), 166.625),
|
||||
id="iq2_xs",
|
||||
),
|
||||
)
|
||||
@@ -90,16 +84,17 @@ def _grid(fmt: _IqFormat, zero: bool = False) -> torch.Tensor:
|
||||
|
||||
|
||||
def _pack(fmt: _IqFormat, extension, weight, grid, scales=None):
|
||||
pack = getattr(extension, fmt.packer)
|
||||
if not fmt.needs_scales:
|
||||
return extension.pack(weight, grid)
|
||||
return pack(weight, grid)
|
||||
if scales is None:
|
||||
scales = torch.zeros(weight.numel() // 256, device=weight.device, dtype=torch.float16)
|
||||
return extension.pack(weight, grid, scales)
|
||||
return pack(weight, grid, scales)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_zero_block_layout(fmt):
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
weight = torch.zeros((2, 256), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
packed = _pack(fmt, extension, weight, _grid(fmt, zero=True))
|
||||
@@ -111,7 +106,7 @@ def test_cuda_ext_iq_zero_block_layout(fmt):
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_encodes_non_zero_block(fmt):
|
||||
"""Exercise the encode loop itself: search, reductions, and the payload writes."""
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
weight = torch.randn((2, 256), device="cuda", dtype=torch.bfloat16, generator=_generator())
|
||||
scales = (weight.float().abs().amax(dim=-1) / fmt.native_max).half()
|
||||
|
||||
@@ -132,7 +127,7 @@ def test_cuda_ext_iq_encodes_non_zero_block(fmt):
|
||||
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_rejects_unsupported_dtype(fmt):
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
weight = torch.ones((1, 256), device="cuda").to(torch.float8_e4m3fn)
|
||||
|
||||
with pytest.raises(RuntimeError, match="supports float32, float64, float16, and bfloat16"):
|
||||
@@ -141,7 +136,7 @@ def test_cuda_ext_iq_rejects_unsupported_dtype(fmt):
|
||||
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_rejects_row_straddling_input(fmt):
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
weight = torch.ones((512, 384), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
with pytest.raises(RuntimeError, match="innermost dimension must be a multiple of 256"):
|
||||
@@ -151,24 +146,24 @@ def test_cuda_ext_iq_rejects_row_straddling_input(fmt):
|
||||
@pytest.mark.parametrize("bad", [float("nan"), float("inf"), float("-inf"), -1.0, -1e-4])
|
||||
def test_cuda_ext_iq2_xs_rejects_invalid_scales(bad):
|
||||
"""A non-finite scale decodes to garbage; a negative one inverts every decoded element."""
|
||||
extension = ext.get_cuda_ext_iq2_xs(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
fmt = _IQ_EXTENSIONS[1].values[0]
|
||||
weight = torch.ones((1, 256), device="cuda", dtype=torch.bfloat16)
|
||||
scales = torch.full((1,), bad, device="cuda", dtype=torch.float16)
|
||||
|
||||
with pytest.raises(RuntimeError, match="scales must be finite and non-negative"):
|
||||
extension.pack(weight, _grid(fmt), scales)
|
||||
extension.iq2_xs_pack(weight, _grid(fmt), scales)
|
||||
|
||||
|
||||
def test_cuda_ext_iq2_xs_negative_zero_scale_packs_as_zero():
|
||||
"""Negative zero is a zero scale: it must take the zero-payload branch, not search."""
|
||||
extension = ext.get_cuda_ext_iq2_xs(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
fmt = _IQ_EXTENSIONS[1].values[0]
|
||||
weight = torch.randn((2, 256), device="cuda", dtype=torch.bfloat16, generator=_generator())
|
||||
grid = _random_grid(fmt)
|
||||
scales = torch.tensor([-0.0, 0.0], device="cuda", dtype=torch.float16)
|
||||
|
||||
packed = extension.pack(weight, grid, scales)
|
||||
packed = extension.iq2_xs_pack(weight, grid, scales)
|
||||
|
||||
assert not packed.any()
|
||||
|
||||
@@ -260,7 +255,7 @@ def test_cuda_ext_iq_encoding_is_optimal(fmt):
|
||||
misplaced index, local scale, delta sign, or sign bit makes the reconstruction worse than
|
||||
the brute-force optimum rather than merely different.
|
||||
"""
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
weight = torch.randn((4, 256), device="cuda", dtype=torch.float32, generator=_generator())
|
||||
grid = _random_grid(fmt)
|
||||
scales = (weight.abs().amax(dim=-1) / fmt.native_max).half() if fmt.needs_scales else None
|
||||
@@ -284,7 +279,7 @@ def test_cuda_ext_iq_encoding_is_optimal(fmt):
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_input_dtype_equivalence(fmt):
|
||||
"""Every accepted input dtype carrying identical values must pack to identical bytes."""
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
# Multiples of 1/16 in [-4, 4) are exact in float16 and bfloat16 as well as the wider types.
|
||||
weight = torch.randint(-64, 64, (2, 256), device="cuda", generator=_generator()).float() / 16
|
||||
grid = _random_grid(fmt)
|
||||
@@ -302,7 +297,7 @@ def test_cuda_ext_iq_input_dtype_equivalence(fmt):
|
||||
@pytest.mark.parametrize("fmt", _IQ_EXTENSIONS)
|
||||
def test_cuda_ext_iq_non_finite_inputs_are_zeroed(fmt):
|
||||
"""NaN and infinity pack as zeros; finite values too large for float32 saturate instead."""
|
||||
extension = fmt.get_extension(raise_if_failed=True)
|
||||
extension = ext.get_cuda_ext_ggml(raise_if_failed=True)
|
||||
clean = torch.randn((2, 256), device="cuda", dtype=torch.float32, generator=_generator())
|
||||
grid = _random_grid(fmt)
|
||||
scales = (clean.abs().amax(dim=-1) / fmt.native_max).half() if fmt.needs_scales else None
|
||||
|
||||
Reference in New Issue
Block a user