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