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
Download code/kernels/test_kernels.py from cklxx/laya-browser: direct link, hf CLI and curl.
- Browser
- Download file 4.59 kB
-
https://huggingface.co/cklxx/laya-browser/resolve/main/code/kernels/test_kernels.py
- Command line
-
hf download hf://cklxx/laya-browser/code/kernels/test_kernels.py
-
curl -L -o test_kernels.py https://huggingface.co/cklxx/laya-browser/resolve/main/code/kernels/test_kernels.py
4.59 kB
| import torch, math, time, sys | |
| sys.path.insert(0, "kernels") | |
| import tl_kernels as K | |
| torch.manual_seed(0) | |
| dev = "cuda" | |
| def err(a, b): return (a.float() - b.float()).abs().max().item() | |
| def bench(f, n=50): | |
| 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 | |
| M, Kd = 8 * 128, 768 | |
| A = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16) | |
| # gemm variants | |
| for N, bias, act in [(2304, False, "none"), (768, True, "gelu"), (3072, True, "relu")]: | |
| W = torch.randn(N, Kd, device=dev, dtype=torch.bfloat16) * 0.02 | |
| b = torch.randn(N, device=dev) | |
| C = torch.empty(M, N, device=dev, dtype=torch.bfloat16) | |
| k = K.gemm_kernel(N, Kd, bias=bias, act=act) | |
| k(A, W, b, C) | |
| ref = A.float() @ W.float().T + (b if bias else 0) | |
| ref = {"none": ref, "gelu": torch.nn.functional.gelu(ref), "relu": torch.relu(ref)}[act] | |
| print(f"gemm N={N} bias={bias} act={act}: maxerr={err(C, ref):.4f} tl={bench(lambda: k(A, W, b, C)):.3f}ms torch={bench(lambda: torch.nn.functional.linear(A, W, b.bfloat16() if bias else None)):.3f}ms") | |
| # geglu | |
| F = 1152; Wi = torch.randn(2 * F, Kd, device=dev, dtype=torch.bfloat16) * 0.02 | |
| C = torch.empty(M, F, device=dev, dtype=torch.bfloat16); k = K.gemm_geglu_kernel(F, Kd); k(A, Wi, C) | |
| x = (A.float() @ Wi.float().T); ref = torch.nn.functional.gelu(x[:, :F]) * x[:, F:] | |
| def tref(): | |
| i, g = torch.nn.functional.linear(A, Wi).chunk(2, -1); return torch.nn.functional.gelu(i) * g | |
| print(f"geglu: maxerr={err(C, ref):.4f} (ref scale {ref.abs().max():.2f}) tl={bench(lambda: k(A, Wi, C)):.3f}ms torch={bench(tref):.3f}ms") | |
| # add_ln | |
| X = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16); R = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16) | |
| w = torch.rand(Kd, device=dev) + 0.5; bb = torch.randn(Kd, device=dev) | |
| for residual, bias in [(True, False), (False, False), (True, True)]: | |
| X2 = X.clone(); Y = torch.empty_like(X) | |
| k = K.add_ln_kernel(Kd, residual=residual, bias=bias); k(X2, R, w, bb, Y) | |
| xr = (X.float() + R.float()) if residual else X.float() | |
| xr_b = xr.bfloat16().float() if residual else xr # kernel stores the residual stream in bf16 | |
| ref = torch.nn.functional.layer_norm(xr_b, (Kd,), w, bb if bias else None, 1e-5) | |
| print(f"add_ln residual={residual} bias={bias}: maxerr Y={err(Y, ref):.4f} X={err(X2, xr):.4f} tl={bench(lambda: k(X2, R, w, bb, Y)):.3f}ms torch={bench(lambda: torch.nn.functional.layer_norm(X2 + R, (Kd,), w.bfloat16(), None, 1e-5)):.3f}ms") | |
| # rope | |
| H, Dh, L = 12, 64, 128; B = M // L | |
| qkv = torch.randn(M, 3 * H * Dh, device=dev, dtype=torch.bfloat16) | |
| inv = 1.0 / (10000 ** (torch.arange(0, Dh, 2, device=dev).float() / Dh)); pos = torch.arange(L, device=dev).float() | |
| fr = torch.outer(pos, inv); cos, sin = fr.cos().contiguous(), fr.sin().contiguous() | |
| q2 = qkv.clone(); k = K.rope_kernel(H, Dh, L); k(q2, cos, sin) | |
| def rot(x): # x [B,L,H,Dh] | |
| c = torch.cat([cos, cos], -1)[None, :, None]; s = torch.cat([sin, sin], -1)[None, :, None] | |
| x1, x2 = x[..., :Dh // 2], x[..., Dh // 2:] | |
| return x * c + torch.cat([-x2, x1], -1) * s | |
| v = qkv.float().view(B, L, 3, H, Dh); ref = v.clone(); ref[:, :, 0] = rot(v[:, :, 0]); ref[:, :, 1] = rot(v[:, :, 1]) | |
| print(f"rope: maxerr={err(q2.view(B, L, 3, H, Dh), ref):.4f} tl={bench(lambda: k(q2, cos, sin)):.3f}ms") | |
| # attention | |
| for L, window in [(128, 0), (128, 64), (1024, 0), (1024, 64)]: | |
| B = 4; qkv = torch.randn(B, L, 3, H, Dh, device=dev, dtype=torch.bfloat16) | |
| lens = torch.tensor([L, L - 5, L // 2 + 3, 7], device=dev, dtype=torch.int32) | |
| O = torch.empty(B, L, H * Dh, device=dev, dtype=torch.bfloat16) | |
| k = K.attn_kernel(B, L, H, Dh, window=window); k(qkv, lens, O) | |
| q, kk, vv = [qkv[:, :, i].transpose(1, 2).float() for i in range(3)] | |
| idx = torch.arange(L, device=dev) | |
| mask = (idx[None, :] < lens[:, None])[:, None, None, :].expand(B, 1, L, L) | |
| if window: mask = mask & ((idx[:, None] - idx[None, :]).abs() <= window)[None, None] | |
| ref = torch.nn.functional.scaled_dot_product_attention(q, kk, vv, attn_mask=mask).transpose(1, 2).reshape(B, L, -1) | |
| valid = (idx[None, :] < lens[:, None]) | |
| e = (O.float() - ref)[valid].abs().max().item() | |
| def sref(): return torch.nn.functional.scaled_dot_product_attention(qkv[:, :, 0].transpose(1, 2), qkv[:, :, 1].transpose(1, 2), qkv[:, :, 2].transpose(1, 2), attn_mask=mask) | |
| print(f"attn L={L} window={window}: maxerr(valid rows)={e:.4f} finite={torch.isfinite(O).all().item()} tl={bench(lambda: k(qkv, lens, O)):.3f}ms torch-sdpa={bench(sref):.3f}ms") | |