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
File size: 11,533 Bytes
adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 4219958 adf912b 8086c56 adf912b 8086c56 | 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 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 | """GPU fast path for laya: TileLang fused kernels + bf16 resident weights + CUDA graphs.
agent = laya.load("convaiinnovations/laya", fast=True) # or agent.accelerate()
Requires CUDA and `pip install laya[fast]` (tilelang). Falls back to the stock forward otherwise.
"""
import sys
import torch
import tl_kernels as K
BF = torch.bfloat16
def _bucket_n(n):
return 1 << max(0, (n - 1).bit_length())
class FastLaya:
def __init__(self, model, max_len=1024, use_graphs=True, verbose=False):
self.m = model
enc = model.encoder
cfg = enc.config
dev = next(model.parameters()).device
self.dev = dev
self.use_graphs = use_graphs
self.verbose = verbose
self.H, self.Dh, self.D = cfg.num_attention_heads, cfg.hidden_size // cfg.num_attention_heads, cfg.hidden_size
self.F = cfg.intermediate_size
self.eps = cfg.norm_eps
self.max_len = max_len
f32 = lambda t: t.detach().float().contiguous()
b16 = lambda t: t.detach().to(BF).contiguous()
zeros = torch.zeros(self.D, device=dev)
self.zeros = {self.D: zeros, 3 * self.D: torch.zeros(3 * self.D, device=dev), 4 * self.D: torch.zeros(4 * self.D, device=dev),
2 * self.F: torch.zeros(2 * self.F, device=dev)}
# --- encoder weights
# embeddings are gathered from an exact fp16 copy of the checkpoint values and upcast to fp32, like the stock path
self.emb_w = enc.embeddings.tok_embeddings.weight.detach().to(torch.float16).contiguous()
self.emb_ln = f32(enc.embeddings.norm.weight)
# HF's bidirectional sliding mask keeps keys with |i - j| <= config.sliding_window (= local_attention // 2);
# ModernBertAttention.sliding_window is that value + 1 (flash-attn's inclusive convention) and must NOT be used here.
win = getattr(cfg, "sliding_window", None) or cfg.local_attention // 2
self.layers = []
for i, lyr in enumerate(enc.layers):
self.layers.append(dict(
attn_ln=None if i == 0 else f32(lyr.attn_norm.weight),
wqkv=b16(lyr.attn.Wqkv.weight), wo=b16(lyr.attn.Wo.weight),
mlp_ln=f32(lyr.mlp_norm.weight), wi=b16(lyr.mlp.Wi.weight), wo2=b16(lyr.mlp.Wo.weight),
window=(win if lyr.attention_type == "sliding_attention" else 0), ltype=lyr.attention_type))
self.final_ln = f32(enc.final_norm.weight)
# --- rotary tables (rounded through bf16 exactly like HF does before applying)
rot = enc.rotary_emb
pos = torch.arange(max_len, device=dev).float()
self.rope = {}
for lt in set(cfg.layer_types):
inv = getattr(rot, f"{lt}_inv_freq").float()
scl = getattr(rot, f"{lt}_attention_scaling")
fr = torch.outer(pos, inv)
self.rope[lt] = ((fr.cos() * scl).to(BF).float().contiguous(), (fr.sin() * scl).to(BF).float().contiguous())
# --- decision head (nn.TransformerEncoderLayer, norm_first, relu)
self.type_emb = b16(model.type_emb.weight)
self.head = []
for lyr in model.head.layers:
sa = lyr.self_attn
self.head.append(dict(
n1w=f32(lyr.norm1.weight), n1b=f32(lyr.norm1.bias), n2w=f32(lyr.norm2.weight), n2b=f32(lyr.norm2.bias),
in_w=b16(sa.in_proj_weight), in_b=f32(sa.in_proj_bias), out_w=b16(sa.out_proj.weight), out_b=f32(sa.out_proj.bias),
l1w=b16(lyr.linear1.weight), l1b=f32(lyr.linear1.bias), l2w=b16(lyr.linear2.weight), l2b=f32(lyr.linear2.bias)))
# --- kernels (M is dynamic, so these compile once)
D, F = self.D, self.F
self.k_qkv = K.gemm_kernel(3 * D, D)
self.k_o = K.gemm_kernel(D, D)
self.k_geglu = K.gemm_geglu_kernel(F, D)
self.k_o2 = K.gemm_kernel(D, F)
self.k_addln = K.add_ln_kernel(D, residual=True, bias=False, eps=self.eps)
self.k_addln_b = K.add_ln_kernel(D, residual=True, bias=True, eps=1e-5)
self.k_ln_b = K.add_ln_kernel(D, residual=False, bias=True, eps=1e-5)
self.k_in = K.gemm_kernel(3 * D, D, bias=True)
self.k_out = K.gemm_kernel(D, D, bias=True)
self.k_ffn1 = K.gemm_kernel(4 * D, D, bias=True, act="relu")
self.k_ffn2 = K.gemm_kernel(D, 4 * D, bias=True)
self._rope_k, self._rope_tab, self._attn_k = None, {}, {}
self.graphs = {}
# ------------------------------------------------------------------ kernels per shape
DYNAMIC_MAX_L = 256 # up to here one dynamic-shape attention kernel is as fast as a static one
LONG_BUCKET = 64 # beyond it, static kernels per (B, L) bucket of this size
def rope_k(self):
if self._rope_k is None:
self._rope_k = K.rope_kernel(self.H, self.Dh)
return self._rope_k
def rope_tab(self, ltype, L):
key = (ltype, L)
if key not in self._rope_tab:
cos, sin = self.rope[ltype]
self._rope_tab[key] = (cos[:L].contiguous(), sin[:L].contiguous())
return self._rope_tab[key]
def attn_k(self, B, L, window):
key = (None, None, window) if L <= self.DYNAMIC_MAX_L else (B, L, window)
if key not in self._attn_k:
self._attn_k[key] = K.attn_kernel(key[0], key[1], self.H, self.Dh, window=window)
return self._attn_k[key]
# ------------------------------------------------------------------ encoder + head on padded [B, L]
def _encode(self, ids, lens, qtype):
"""ids [B,L] long (padded), lens [B] int32, qtype [B] long -> hidden [B, L, D] bf16"""
B, L = ids.shape
M, D = B * L, self.D
dev = self.dev
emb = torch.nn.functional.embedding(ids, self.emb_w).view(M, D).float()
# residual stream = embeddings.norm(emb), kept in fp32 exactly like the stock autocast path
X = torch.nn.functional.layer_norm(emb, (D,), self.emb_ln, None, self.eps)
Y = X.to(BF) # layer 0 attends to it directly (attn_norm = Identity)
qkv = torch.empty(M, 3 * D, device=dev, dtype=BF)
O = torch.empty(M, D, device=dev, dtype=BF)
G = torch.empty(M, self.F, device=dev, dtype=BF)
z = self.zeros
nl = len(self.layers)
for i, ly in enumerate(self.layers):
self.k_qkv(Y, ly["wqkv"], z[3 * D], qkv) # Y = attn_norm(X) (layer 0: X itself)
cos, sin = self.rope_tab(ly["ltype"], L)
self.rope_k()(qkv, cos, sin)
self.attn_k(B, L, ly["window"])(qkv.view(B, L, 3, self.H, self.Dh), lens, O.view(B, L, D))
self.k_o(O, ly["wo"], z[D], Y) # Y = attn out
self.k_addln(X, Y, ly["mlp_ln"], z[D], Y) # X += Y ; Y = mlp_norm(X)
self.k_geglu(Y, ly["wi"], G)
self.k_o2(G, ly["wo2"], z[D], Y) # Y = mlp out
nxt = self.layers[i + 1]["attn_ln"] if i + 1 < nl else self.final_ln
self.k_addln(X, Y, nxt, z[D], Y) # X += Y ; Y = next norm(X)
# decision head: h = final_norm(x) + type_emb ; 2 x pre-norm transformer layers (relu ffn)
X = (Y.view(B, L, D).float() + self.type_emb[qtype].float()[:, None, :]).view(M, D).contiguous() # fp32 stream for the head
F1 = torch.empty(M, 4 * D, device=dev, dtype=BF)
for j, h in enumerate(self.head):
self.k_ln_b(X, Y, h["n1w"], h["n1b"], Y)
self.k_in(Y, h["in_w"], h["in_b"], qkv)
self.attn_k(B, L, 0)(qkv.view(B, L, 3, self.H, self.Dh), lens, O.view(B, L, D))
self.k_out(O, h["out_w"], h["out_b"], Y)
self.k_addln_b(X, Y, h["n2w"], h["n2b"], Y) # X += attn ; Y = norm2(X)
self.k_ffn1(Y, h["l1w"], h["l1b"], F1)
self.k_ffn2(F1, h["l2w"], h["l2b"], Y)
X = X + Y.float() # residual (fp32, torch, last op)
return X.view(B, L, D)
def _encode_graphed(self, ids, lens, qtype):
key = tuple(ids.shape)
g = self.graphs.get(key)
if g is None:
s_ids, s_lens, s_q = ids.clone(), lens.clone(), qtype.clone()
st = torch.cuda.Stream()
st.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(st):
for _ in range(2):
self._encode(s_ids, s_lens, s_q) # warm-up (compiles kernels, allocs)
torch.cuda.current_stream().wait_stream(st)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
s_out = self._encode(s_ids, s_lens, s_q)
g = self.graphs[key] = (graph, s_ids, s_lens, s_q, s_out)
if self.verbose:
print(f"[fast_laya] captured CUDA graph for shape {key}", file=sys.stderr)
graph, s_ids, s_lens, s_q, s_out = g
s_ids.copy_(ids); s_lens.copy_(lens); s_q.copy_(qtype)
graph.replay()
return s_out
# ------------------------------------------------------------------ DecisionModel.forward replacement
@torch.no_grad()
def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, detach_encoder=False):
m = self.m
N, L0 = input_ids.shape
g = 16 if L0 <= self.DYNAMIC_MAX_L else self.LONG_BUCKET
L = min(self.max_len, ((L0 + g - 1) // g) * g)
B = _bucket_n(N)
ids = torch.zeros(B, L, dtype=torch.long, device=self.dev)
ids[:N, :L0] = input_ids
lens = torch.zeros(B, dtype=torch.int32, device=self.dev)
lens[:N] = attention_mask.sum(1).to(torch.int32)
qt = torch.zeros(B, dtype=torch.long, device=self.dev)
qt[:N] = qtype
h = (self._encode_graphed if self.use_graphs else self._encode)(ids, lens, qt)
h = h[:N, :L0].float()
# ---- scorer / act head (tiny; identical to laya.common.DecisionModel.forward)
idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1))
mk = torch.gather(h, 1, idx)
logits = m.scorer(mk).squeeze(-1).float()
logits = logits.masked_fill(~marker_mask, -1e4)
p = torch.softmax(logits, -1)
k = marker_mask.sum(-1).clamp(min=2).float()
ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k)
# a choice with a single option (e.g. only WAIT left after exclusions): pad the second slot, as upstream does
top2 = p.topk(2, -1).values if p.shape[-1] >= 2 else torch.stack([p[:, 0], torch.zeros_like(p[:, 0])], dim=-1)
feats = torch.stack([top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], -1)
pooled = h[:, 0].float()
act_logits = m.act_head(torch.cat([pooled, feats], -1))
return logits, act_logits
def accelerate(agent, use_graphs=True, verbose=False):
"""Patch a laya Agent in place so agent.predict() uses the TileLang fast path. Returns the FastLaya."""
fast = FastLaya(agent.model, max_len=agent.cfg.get("max_len", 1024), use_graphs=use_graphs, verbose=verbose)
if not hasattr(agent, "_orig_forward"):
agent._orig_forward = agent.model.forward
agent.model.forward = fast.forward
agent.fast = fast
return fast
def restore(agent):
if hasattr(agent, "_orig_forward"):
agent.model.forward = agent._orig_forward
del agent._orig_forward
|