Feature Extraction
Transformers
Safetensors
Laya
English
multilingual
laya_browser
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)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("cklxx/laya-browser", trust_remote_code=True, device_map="auto") - Laya
How to use cklxx/laya-browser with Laya:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
File size: 3,337 Bytes
454b3e6 | 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 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 | """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")
|