wikiqwen / fp8_kernel.py
devon7y's picture
WikiQwen: app, requirements, README
5ef6851 verified
Raw History Blame Contribute Delete
1.59 kB
"""Make transformers 5.14.x run block-wise FP8 checkpoints without the `kernels` package: download the pinned
kernels-community/finegrained-fp8 build (v4) and plug it in, the same way scripts/hpc/fp8_35b.py does offline."""
import importlib.util, os, sys
REPO, REVISION = "kernels-community/finegrained-fp8", "v4"
def ensure():
from huggingface_hub import snapshot_download
from transformers.integrations import finegrained_fp8 as F
root = snapshot_download(REPO, revision=REVISION, allow_patterns=["build/torch-cuda/*"])
build = os.path.join(root, "build", "torch-cuda")
name = "_wikiqwen_finegrained_fp8"
spec = importlib.util.spec_from_file_location(name, os.path.join(build, "__init__.py"), submodule_search_locations=[build])
mod = importlib.util.module_from_spec(spec)
sys.modules[name] = mod
spec.loader.exec_module(mod)
missing = [s for s in ("matmul_2d", "matmul_batched", "matmul_grouped") if not hasattr(mod, s)]
if missing:
raise ImportError(f"FP8 kernel build {build} lacks {missing}")
fg = F.FineGrainedFP8(matmul=mod.matmul_2d, batched_matmul=mod.matmul_batched, grouped_matmul=mod.matmul_grouped)
F._load_finegrained_fp8_kernel = lambda: fg # 5.14.x looks this name up at call time
try:
from transformers.integrations import hub_kernels
hub_kernels._KERNEL_MODULE_MAPPING["finegrained-fp8"] = mod
except Exception:
pass
if hasattr(F, "_FINEGRAINED_FP8"): # transformers >= 5.17
F._FINEGRAINED_FP8 = fg
print("fp8 kernel ready:", build, flush=True)
return fg