Feature Extraction
Transformers
Safetensors
English
multilingual
laya_browser
laya
custom_code
system-1
browser-agent
web-navigation
decision-model
mmbert
mind2web
tilelang
Instructions to use cklxx/laya-browser with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use cklxx/laya-browser with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="cklxx/laya-browser", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("cklxx/laya-browser", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download code/kernels/tune.py from cklxx/laya-browser: direct link, hf CLI and curl.
- Browser
- Download file 3.34 kB
-
https://huggingface.co/cklxx/laya-browser/resolve/main/code/kernels/tune.py
- Command line
-
hf download hf://cklxx/laya-browser/code/kernels/tune.py
-
curl -L -o tune.py https://huggingface.co/cklxx/laya-browser/resolve/main/code/kernels/tune.py
3.34 kB
| """Sweep tile configs for the GEMM / GEGLU / attention kernels at the shapes laya actually hits.""" | |
| import torch, sys, time, itertools, json | |
| sys.path.insert(0, "kernels"); import tl_kernels as K | |
| dev = "cuda" | |
| def bench(f, n=20): | |
| for _ in range(3): f() | |
| torch.cuda.synchronize(); t = time.perf_counter() | |
| for _ in range(n): f() | |
| torch.cuda.synchronize(); return (time.perf_counter() - t) / n * 1e3 | |
| res = {"gemm": {}, "geglu": {}, "attn": {}} | |
| Ms = [128, 1024, 8192, 28672] | |
| cfgs = [(bm, bn, th) for bm in (64, 128) for bn in (64, 128, 256) for th in (128, 256)] | |
| for (N, Kd) in [(2304, 768), (768, 768), (768, 1152)]: | |
| W = torch.randn(N, Kd, device=dev, dtype=torch.bfloat16) * 0.02; b = torch.zeros(N, device=dev) | |
| for M in Ms: | |
| A = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16); C = torch.empty(M, N, device=dev, dtype=torch.bfloat16) | |
| tref = bench(lambda: torch.nn.functional.linear(A, W)) | |
| best = None | |
| for bm, bn, th in cfgs: | |
| try: | |
| k = K.gemm_kernel(N, Kd, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, W, b, C)) | |
| if best is None or ms < best[0]: best = (ms, (bm, bn, th)) | |
| except Exception as e: print("fail", N, Kd, M, (bm, bn, th), str(e)[:60], flush=True) | |
| tf = 2 * M * N * Kd / best[0] / 1e9 | |
| print(f"gemm N={N} K={Kd} M={M}: best {best[1]} {best[0]:.3f}ms ({tf:.1f} TFLOPS) cublas {tref:.3f}ms (default 64x128x128thr)", flush=True) | |
| res["gemm"][f"{N},{Kd},{M}"] = best | |
| for M in Ms: | |
| A = torch.randn(M, 768, device=dev, dtype=torch.bfloat16); Wi = torch.randn(2304, 768, device=dev, dtype=torch.bfloat16) * 0.02; C = torch.empty(M, 1152, device=dev, dtype=torch.bfloat16) | |
| best = None | |
| for bm, bn, th in cfgs: | |
| if bn > 128: continue | |
| try: | |
| k = K.gemm_geglu_kernel(1152, 768, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, Wi, C)) | |
| if best is None or ms < best[0]: best = (ms, (bm, bn, th)) | |
| except Exception as e: print("fail geglu", M, (bm, bn, th), str(e)[:60], flush=True) | |
| print(f"geglu M={M}: best {best[1]} {best[0]:.3f}ms ({2*M*2304*768/best[0]/1e9:.1f} TFLOPS)", flush=True) | |
| res["geglu"][str(M)] = best | |
| H, Dh = 12, 64 | |
| for (B, L) in [(32, 1024), (32, 128), (4, 128)]: | |
| qkv = torch.randn(B, L, 3, H, Dh, device=dev, dtype=torch.bfloat16); lens = torch.full((B,), L - 3, device=dev, dtype=torch.int32); O = torch.empty(B, L, H * Dh, device=dev, dtype=torch.bfloat16) | |
| for window in (0, 65): | |
| best = None | |
| for bm, bn, st, th in itertools.product((64, 128), (64, 128), (1, 2), (128, 256)): | |
| try: | |
| k = K.attn_kernel(B, L, H, Dh, window=window, bm=bm, bn=bn, stages=st, threads=th); ms = bench(lambda: k(qkv, lens, O)) | |
| if best is None or ms < best[0]: best = (ms, (bm, bn, st, th)) | |
| except Exception as e: print("fail attn", (B, L, window), (bm, bn, st, th), str(e)[:60], flush=True) | |
| fl = 4 * B * H * L * L * Dh if window == 0 else 4 * B * H * L * (2 * 65 + 1) * Dh | |
| print(f"attn B={B} L={L} window={window}: best {best[1]} {best[0]:.3f}ms ({fl/best[0]/1e9:.1f} TFLOPS)", flush=True) | |
| res["attn"][f"{B},{L},{window}"] = best | |
| json.dump(res, open("kernels/tune_results.json", "w"), indent=1) | |
| print("DONE") | |