Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion bonsai/models/efficientnet/params.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,7 @@ def load_pretrained_weights(model: model_lib.EfficientNet, pretrained_weights: d
f"pre-trained weight has {weight_np.shape}."
)

param_to_update.value = jnp.array(weight_np)
param_to_update.set_value(jnp.array(weight_np))

return model

Expand Down
28 changes: 14 additions & 14 deletions bonsai/models/qwen3/modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -218,7 +218,7 @@ def __init__(self, einsum_str: str, shape: tuple[int, ...], *, shd: ShardingSpec

@jax.named_scope("einsum")
def __call__(self, x: ArrayLike) -> Array:
return jnp.einsum(self.einsum_str, x, self.w.value)
return jnp.einsum(self.einsum_str, x, self.w[...])


def _generate_pos_embeddings(
Expand Down Expand Up @@ -250,7 +250,7 @@ def __init__(self, dim: int, cfg: ModelConfig, *, rngs: nnx.Rngs):
def __call__(self, x: Array) -> Array:
dtype = x.dtype
rms = jnp.sqrt(jnp.mean(jnp.astype(x, jnp.float32) ** 2, axis=-1, keepdims=True) + self.norm_eps)
return jnp.astype(self.scale.value * x / rms, dtype)
return jnp.astype(self.scale[...] * x / rms, dtype)


def count_left_pads(x: jax.Array) -> int:
Expand Down Expand Up @@ -300,28 +300,28 @@ def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Ar
# RoPE and Cache Logic
left_pads = count_left_pads(segment_ids)
left_pads = shard(left_pads, P(self.shd_cfg.act_btnh[0]))
cache.start_ind.value = jnp.where(cache.start_ind.value < 0, left_pads, cache.start_ind.value)
position_ids = compute_positions_from_segment_ids(segment_ids) + cache.cur_ind.value
cache.start_ind.set_value(jnp.where(cache.start_ind[...] < 0, left_pads, cache.start_ind[...]))
position_ids = compute_positions_from_segment_ids(segment_ids) + cache.cur_ind[...]
sin, cos = _generate_pos_embeddings(position_ids, self.head_dim)
query_proj = apply_rope(query_proj, sin, cos)
key_proj = apply_rope(key_proj, sin, cos)

# Update K/V cache [B, S, K, H]
slice_indices = (0, cache.cur_ind.value, 0, 0)
cache.v_cache.value = jax.lax.dynamic_update_slice(cache.v_cache.value, value_proj, slice_indices)
cache.k_cache.value = jax.lax.dynamic_update_slice(cache.k_cache.value, key_proj, slice_indices)
slice_indices = (0, cache.cur_ind[...], 0, 0)
cache.v_cache[...] = jax.lax.dynamic_update_slice(cache.v_cache[...], value_proj, slice_indices)
cache.k_cache[...] = jax.lax.dynamic_update_slice(cache.k_cache[...], key_proj, slice_indices)

b, t, n, h = query_proj.shape

# GQA reshape and attention logits
query_proj_gqa = query_proj.reshape((b, t, self.num_kv_heads, self.n_rep, h))
attn_logits = jnp.einsum("BTKGH,BSKH->BTSKG", query_proj_gqa, cache.k_cache.value) * self.scale
attn_logits = jnp.einsum("BTKGH,BSKH->BTSKG", query_proj_gqa, cache.k_cache[...]) * self.scale

# Masking and Softmax
q_pos = cache.cur_ind.value + jnp.arange(t, dtype=jnp.int32)[None, :] - cache.start_ind.value[:, None]
q_pos = cache.cur_ind[...] + jnp.arange(t, dtype=jnp.int32)[None, :] - cache.start_ind[:, None]
ts = jnp.arange(cache.size, dtype=jnp.int32) # (cache.size,)
kv_segment_ids = (ts[None, :] >= cache.start_ind.value[:, None]) & (ts[None, :] < cache.cur_ind.value + t)
k_pos = ts[None, :] - cache.start_ind.value[:, None] # (b, cache.size)
kv_segment_ids = (ts[None, :] >= cache.start_ind[:, None]) & (ts[None, :] < cache.cur_ind[...] + t)
k_pos = ts[None, :] - cache.start_ind[:, None] # (b, cache.size)
causal_mask = k_pos[:, None, :] <= q_pos[:, :, None]
segment_mask = kv_segment_ids[:, None, :] == segment_ids[:, :, None]
final_mask = causal_mask & segment_mask # (B, T, S)
Expand All @@ -330,10 +330,10 @@ def __call__(self, x: Array, cache: LayerCache | None, segment_ids: Array) -> Ar

# Softmax
attn_weights = jax.nn.softmax(attn_logits.astype(jnp.float32), axis=2).astype(attn_logits.dtype)
qkv = jnp.einsum("BTSKG,BSKH->BTKGH", attn_weights, cache.v_cache.value)
qkv = jnp.einsum("BTSKG,BSKH->BTKGH", attn_weights, cache.v_cache[...])
qkv = qkv.reshape((b, t, n, h))

cache.cur_ind.value = cache.cur_ind.value + t
cache.cur_ind[...] = cache.cur_ind[...] + t
return shard(self.o_proj(qkv), self.shd_cfg.act_btd)

@property
Expand Down Expand Up @@ -399,7 +399,7 @@ def init_cache(
return [LayerCache(cfg, batch_size, cache_size, dtype) for _ in range(cfg.num_layers)]

def __call__(self, tokens, segment_ids, cache, num_right_pads):
x = self.embedder.embedding.value.at[(tokens,)].get(out_sharding=self.out_emb_shd)
x = self.embedder.embedding[...].at[(tokens,)].get(out_sharding=self.out_emb_shd)
Comment thread
lingebeng marked this conversation as resolved.
for i, layer in enumerate(self.layers):
x = layer(x, cache[i], segment_ids)
logits = self.lm_head(self.final_norm(x))
Expand Down
4 changes: 2 additions & 2 deletions bonsai/models/qwen3/tests/test_outputs_qwen3.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def _setup_torch_attn(self, input_embeddings: torch.Tensor, attention_mask: None
def _nnx_forward_logits(self, cache: modeling.Cache, tokens: jax.Array, dtype: DTypeLike = jnp.float32):
"""Forward pass for the nnx model"""
segment_ids = 1 * (tokens != self.tokenizer.pad_token_id)
x = self.nnx_model.embedder.embedding.value.at[(tokens,)].get().astype(jnp.float32)
x = self.nnx_model.embedder.embedding[...].at[(tokens,)].get().astype(jnp.float32)
Comment thread
lingebeng marked this conversation as resolved.
for i, layer in enumerate(self.nnx_model.layers):
x = layer(x, cache[i], segment_ids).astype(dtype)
nnx_logits = self.nnx_model.lm_head(self.nnx_model.final_norm(x))
Expand Down Expand Up @@ -133,7 +133,7 @@ def test_embedder(self):
tx = torch.randint(0, self.torch_model.config.vocab_size, size=(self.batch_size, self.num_input_tokens))
jx = jnp.array(tx.cpu().detach().numpy())

jy, ty = nm.embedding.value.at[(jx,)].get(), tm(tx)
jy, ty = nm.embedding[...].at[(jx,)].get(), tm(tx)
Comment thread
lingebeng marked this conversation as resolved.
torch.testing.assert_close(
torch.tensor(np.array(jy, dtype=np.float32)),
ty,
Expand Down
8 changes: 3 additions & 5 deletions bonsai/models/sam2/model_hiera_det.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,12 +293,10 @@ def __init__(

def _get_pos_embed(self, hw: tuple[int, int]) -> jnp.ndarray:
h, w = hw
pos_embed = jax.image.resize(
self.pos_embed.value, shape=(1, self.pos_embed.value.shape[1], h, w), method="bicubic"
)
pos_embed = jax.image.resize(self.pos_embed[...], shape=(1, self.pos_embed.shape[1], h, w), method="bicubic")

tile_factors = [1, h // self.pos_embed_window.value.shape[2], w // self.pos_embed_window.value.shape[3]]
window_embed = jnp.tile(self.pos_embed_window.value, tile_factors)
tile_factors = [1, h // self.pos_embed_window.shape[2], w // self.pos_embed_window.shape[3]]
window_embed = jnp.tile(self.pos_embed_window[...], tile_factors)
pos_embed = pos_embed + window_embed
return jnp.transpose(pos_embed, (0, 2, 3, 1)) # BCHW -> BHWC

Expand Down
4 changes: 2 additions & 2 deletions bonsai/models/vit/modeling.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,7 @@ def __call__(self, pixel_values: jnp.ndarray, *, rngs: nnx.Rngs | None) -> jnp.n

num_new_patch_patches = h * w

stored_pos_embeddings = self.pos_embeddings.value
stored_pos_embeddings = self.pos_embeddings[...]
num_stored_patches = stored_pos_embeddings.shape[1]

if num_stored_patches == num_new_patch_patches + 1:
Expand All @@ -87,7 +87,7 @@ def __call__(self, pixel_values: jnp.ndarray, *, rngs: nnx.Rngs | None) -> jnp.n
stored_pos_embeddings, num_tokens=num_new_patch_patches + 1, has_class_token=True
)

cls_tokens = jnp.tile(self.cls_token.value, (b, 1, 1))
cls_tokens = jnp.tile(self.cls_token[...], (b, 1, 1))
embeddings = jnp.concatenate((cls_tokens, embeddings), axis=1)

embeddings = embeddings + current_pos_embeddings
Expand Down