File size: 1,591 Bytes
5ef6851 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 | """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
|