adjaxt 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
adjaxt/__init__.py ADDED
@@ -0,0 +1,36 @@
1
+ """adjaxt — distributed training with JAX on free compute, made not sloppy."""
2
+
3
+ from adjaxt.config import (
4
+ RMSNormConfig,
5
+ SwiGLUConfig,
6
+ StandardAttnImplementation,
7
+ GQAAttnConfig,
8
+ Qwen3AttnConfig,
9
+ Qwen3MLPConfig,
10
+ Qwen3MoEBlockConfig,
11
+ Qwen3MoELayerConfig,
12
+ Qwen3MoEModelConfig,
13
+ compute_rope_freqs,
14
+ )
15
+ from adjaxt import layers, models, model_maps, sharding
16
+
17
+ __version__ = "0.0.1"
18
+
19
+ __all__ = [
20
+ # Configs
21
+ "RMSNormConfig",
22
+ "SwiGLUConfig",
23
+ "StandardAttnImplementation",
24
+ "GQAAttnConfig",
25
+ "Qwen3AttnConfig",
26
+ "Qwen3MLPConfig",
27
+ "Qwen3MoEBlockConfig",
28
+ "Qwen3MoELayerConfig",
29
+ "Qwen3MoEModelConfig",
30
+ "compute_rope_freqs",
31
+ # Submodules
32
+ "layers",
33
+ "models",
34
+ "model_maps",
35
+ "sharding",
36
+ ]
adjaxt/approx.py ADDED
File without changes
adjaxt/checkpoint.py ADDED
File without changes
adjaxt/config.py ADDED
@@ -0,0 +1,88 @@
1
+ from dataclasses import dataclass, field
2
+ from typing import Dict, Any, Tuple
3
+ from enum import StrEnum
4
+ import jax
5
+ import jax.numpy as jnp
6
+
7
+ @dataclass(frozen=True)
8
+ class RMSNormConfig:
9
+ dim: int
10
+ eps: float = 1e-6
11
+
12
+ @dataclass(frozen=True)
13
+ class SwiGLUConfig:
14
+ in_dim: int
15
+ hidden_dim: int
16
+
17
+ class StandardAttnImplementation(StrEnum):
18
+ CUDNN = "cudnn"
19
+ XLA = "xla"
20
+
21
+ @dataclass(frozen=True)
22
+ class GQAAttnConfig:
23
+ implementation: StandardAttnImplementation
24
+ num_kv_groups: int
25
+ is_causal: bool
26
+
27
+ def compute_rope_freqs(seq_len: int, head_dim: int, theta: float = 10000.0):
28
+ dim_indices = jnp.arange(0, head_dim, 2, dtype=jnp.float32)
29
+ inv_freq = 1.0 / (theta ** (dim_indices / head_dim))
30
+ t = jnp.arange(seq_len, dtype=jnp.float32)
31
+ freqs = jnp.outer(t, inv_freq)
32
+
33
+ freqs = jnp.concatenate([freqs, freqs], axis=-1)
34
+
35
+ cos = jnp.cos(freqs)[None, :, None, :]
36
+ sin = jnp.sin(freqs)[None, :, None, :]
37
+ return cos, sin
38
+
39
+ @dataclass(frozen=True)
40
+ class Qwen3AttnConfig:
41
+ gqa_conf: GQAAttnConfig
42
+ q_rms_conf: RMSNormConfig
43
+ k_rms_conf: RMSNormConfig
44
+ d: int
45
+ head_dim: int
46
+ rope_theta: float
47
+ cos_table: jax.Array = field(init=False, compare=False, hash=False)
48
+ sin_table: jax.Array = field(init=False, compare=False, hash=False)
49
+ max_position_embeddings: int
50
+ n_layers: int
51
+
52
+ def __post_init__(self):
53
+ cos, sin = compute_rope_freqs(
54
+ self.max_position_embeddings,
55
+ self.head_dim,
56
+ self.rope_theta
57
+ )
58
+ object.__setattr__(self, "cos_table", cos)
59
+ object.__setattr__(self, "sin_table", sin)
60
+
61
+ @dataclass(frozen=True)
62
+ class Qwen3MLPConfig:
63
+ act_fn: function
64
+ in_dim: int
65
+ hidden_dim: int
66
+
67
+ @dataclass(frozen=True)
68
+ class Qwen3MoEBlockConfig:
69
+ mlp_conf: Qwen3MLPConfig
70
+ top_k: int
71
+ num_experts: int
72
+ d_model: int
73
+
74
+ @dataclass(frozen=True)
75
+ class Qwen3MoELayerConfig:
76
+ input_rms_conf: RMSNormConfig
77
+ attn_conf: Qwen3AttnConfig
78
+ post_attn_rms_conf: RMSNormConfig
79
+ moe_block_conf: Qwen3MoEBlockConfig
80
+
81
+ @dataclass(frozen=True)
82
+ class Qwen3MoEModelConfig:
83
+ moe_layer_conf: Qwen3MoELayerConfig
84
+ final_rms_conf: RMSNormConfig
85
+ num_decoder_blocks: int
86
+ vocab_size: int
87
+ d_model: int
88
+ tie_word_embeddings: bool = False
adjaxt/data.py ADDED
File without changes
adjaxt/hub.py ADDED
File without changes
adjaxt/init.py ADDED
File without changes
adjaxt/layers.py ADDED
@@ -0,0 +1,322 @@
1
+ import jax, jax.numpy as jnp
2
+ import math
3
+ from adjaxt.config import *
4
+
5
+ #def init_normal_layer(key, shape, std: float) -> jax.Array:
6
+ # return jax.random.normal(key, shape, jnp.bfloat16) * std
7
+
8
+
9
+ #===================================================================================
10
+ #RMS Norm
11
+
12
+ def rms_norm(x, w, cfg) -> jax.Array:
13
+ """
14
+ Args:
15
+ x: [..., D]
16
+ w: [D]
17
+ Returns:
18
+ out: [..., D]
19
+ """
20
+ x32 = x.astype(jnp.float32)
21
+ ms = jnp.mean(jnp.square(x32), axis=-1, keepdims=True)
22
+ return w * (x32 * jax.lax.rsqrt(ms + cfg.eps)).astype(x.dtype)
23
+
24
+ def rms_norm_init(key, cfg: RMSNormConfig) -> jax.Array:
25
+ """
26
+ Returns:
27
+ w: [dim]
28
+ """
29
+ return jnp.ones((cfg.dim,), dtype=jnp.float32)
30
+
31
+ #===================================================================================
32
+ #RoPE
33
+
34
+ def rotate_half(x: jax.Array) -> jax.Array:
35
+ """
36
+ Args:
37
+ x: [B, S, H, D]
38
+ Returns:
39
+ out: [B, S, H, D]
40
+ """
41
+ x1 = x[..., : x.shape[-1] // 2]
42
+ x2 = x[..., x.shape[-1] // 2 :]
43
+ return jnp.concatenate((-x2, x1), axis=-1)
44
+
45
+ def apply_rope(x: jax.Array, cos: jax.Array, sin: jax.Array) -> jax.Array:
46
+ return x * cos + rotate_half(x) * sin
47
+
48
+ def apply_rot_pos_emb(q: jax.Array, k: jax.Array, cos: jax.Array, sin: jax.Array):
49
+ """
50
+ Args:
51
+ q: [B, S, H_q, D]
52
+ k: [B, S, H_kv, D]
53
+ cos: [1, S, 1, D]
54
+ sin: [1, S, 1, D]
55
+ Returns:
56
+ q_rot, k_rot: ([B, S, H_q, D], [B, S, H_kv, D])
57
+ """
58
+ # q: (batch, seq_len, num_heads, head_dim)
59
+ # cos: (1, seq_len, 1, head_dim)
60
+ q_rot = (q * cos) + (rotate_half(q) * sin)
61
+ k_rot = (k * cos) + (rotate_half(k) * sin)
62
+ return q_rot, k_rot
63
+
64
+ #===================================================================================
65
+ #SwiGLU
66
+
67
+ def swiglu(x: jax.Array, w: dict, cfg: SwiGLUConfig) -> jax.Array:
68
+ return (jax.nn.silu(x @ w["w_gate"]) * (x @ w["w_up"])) @ w["w_down"]
69
+
70
+ def swiglu_init(key, cfg: SwiGLUConfig) -> dict:
71
+ k1, k2, k3 = jax.random.split(key, 3)
72
+ def kaiming(k, shape):
73
+ return jax.random.normal(k, shape, dtype=jnp.float32) * jnp.sqrt(2.0 / shape[0])
74
+ return {
75
+ "w_gate": kaiming(k1, (cfg.in_dim, cfg.hidden_dim)),
76
+ "w_up": kaiming(k2, (cfg.in_dim, cfg.hidden_dim)),
77
+ "w_down": kaiming(k3, (cfg.hidden_dim, cfg.in_dim))
78
+ }
79
+
80
+ #===================================================================================
81
+ #RepeatKV
82
+
83
+ def repeat_kv(hidden_states: jax.Array, n_rep: int) -> jax.Array:
84
+ """
85
+ Expands KV heads to match query heads.
86
+ Args:
87
+ x: [B, S, H_kv, D]
88
+ Returns:
89
+ out: [B, S, H_kv * n_rep, D]
90
+ """
91
+ if n_rep == 1:
92
+ return hidden_states
93
+ return jnp.repeat(hidden_states, n_rep, axis=2)
94
+
95
+ #===================================================================================
96
+ #GQAAttention
97
+
98
+ def gqa_attn(
99
+ q: jax.Array,
100
+ k: jax.Array,
101
+ v: jax.Array,
102
+ cfg: GQAAttnConfig,
103
+ attn_mask: jax.Array = None
104
+ ) -> jax.Array:
105
+ """
106
+ Args:
107
+ q: [B, S_q, H_q, D]
108
+ k: [B, S_kv, H_kv, D]
109
+ v: [B, S_kv, H_kv, D]
110
+ attn_mask: [B, H_q, S_q, S_kv] or None
111
+ Returns:
112
+ out: [B, S_q, H_q, D]
113
+ """
114
+
115
+ batch, q_len, num_q_heads, head_dim = q.shape
116
+ _, kv_len, num_kv_heads, _ = k.shape
117
+
118
+ num_groups = num_q_heads // num_kv_heads
119
+
120
+ if num_groups > 1:
121
+ k = jnp.repeat(k, repeats=num_groups, axis=2)
122
+ v = jnp.repeat(v, repeats=num_groups, axis=2)
123
+
124
+ causal_flag = cfg.is_causal if attn_mask is None else False
125
+
126
+ out = jax.nn.dot_product_attention(
127
+ query=q,
128
+ key=k,
129
+ value=v,
130
+ mask=attn_mask,
131
+ is_causal=causal_flag,
132
+ implementation=cfg.implementation,
133
+ )
134
+
135
+ return out
136
+
137
+ #===================================================================================
138
+ #Qwen3Attn
139
+
140
+ def qwen3_attn(
141
+ x: jax.Array,
142
+ w: dict,
143
+ cfg: Qwen3AttnConfig
144
+ ):
145
+ """
146
+ Args:
147
+ x: [B, S, D_model]
148
+ w: Dict containing projection weights and head norms
149
+ Returns:
150
+ out: [B, S, D_model]
151
+ """
152
+ seq_len = x.shape[1]
153
+ input_shape = x.shape[:-1]
154
+ hidden_shape = (*input_shape, -1, cfg.head_dim)
155
+
156
+ # Shapes: (batch, seq_len, num_heads, head_dim)
157
+ q = rms_norm((x @ w["q_proj"]).reshape(hidden_shape), w["q_norm"], cfg.q_rms_conf)
158
+ k = rms_norm((x @ w["k_proj"]).reshape(hidden_shape), w["k_norm"], cfg.k_rms_conf)
159
+ v = (x @ w["v_proj"]).reshape(hidden_shape) # keep (B, S, KV_H, D)
160
+
161
+ # cos/sin shape: (1, seq_len, 1, head_dim)
162
+ cos = cfg.cos_table[:, :seq_len, :, :]
163
+ sin = cfg.sin_table[:, :seq_len, :, :]
164
+ q, k = apply_rot_pos_emb(q, k, cos, sin)
165
+
166
+ attn_out = gqa_attn(
167
+ q,
168
+ k,
169
+ v,
170
+ cfg.gqa_conf
171
+ )
172
+ attn_out = attn_out.reshape(*input_shape, -1)
173
+ attn_out = attn_out @ w["o_proj"]
174
+ return attn_out
175
+
176
+
177
+ def qwen3_attn_init(
178
+ key,
179
+ cfg: Qwen3AttnConfig
180
+ ) -> dict:
181
+ k1, k2, k3, k4, k5, k6 = jax.random.split(key, 6)
182
+ res = 1.0 / math.sqrt(2 * cfg.n_layers)
183
+
184
+ # Calculate dimensions based on GQA heads
185
+ q_dim = cfg.d # num_heads * head_dim
186
+ kv_dim = (cfg.d // cfg.gqa_conf.num_kv_groups) # num_kv_heads * head_dim
187
+
188
+ w = {
189
+ "q_proj": jax.random.normal(k1, (cfg.d, q_dim), jnp.bfloat16) / math.sqrt(cfg.d),
190
+ "k_proj": jax.random.normal(k2, (cfg.d, kv_dim), jnp.bfloat16) / math.sqrt(cfg.d),
191
+ "v_proj": jax.random.normal(k3, (cfg.d, kv_dim), jnp.bfloat16) / math.sqrt(cfg.d),
192
+ "o_proj": jax.random.normal(k4, (q_dim, cfg.d), jnp.bfloat16) * res / math.sqrt(q_dim),
193
+ "q_norm": jnp.ones((cfg.head_dim,), dtype=jnp.bfloat16),
194
+ "k_norm": jnp.ones((cfg.head_dim,), dtype=jnp.bfloat16),
195
+ }
196
+ return w
197
+
198
+ #===================================================================================
199
+ #Qwen3MLP
200
+
201
+ def qwen3_mlp(x: jax.Array, w: dict, cfg: Qwen3MLPConfig) -> jax.Array:
202
+ """
203
+ Args:
204
+ x: [..., D_in]
205
+ w: {"w_gate": [D_in, D_hidden], "w_up": [D_in, D_hidden], "w_down": [D_hidden, D_in]}
206
+ Returns:
207
+ out: [..., D_in]
208
+ """
209
+ return (cfg.act_fn(x @ w["w_gate"]) * (x @ w["w_up"])) @ w["w_down"]
210
+
211
+ def qwen3_mlp_init(key, cfg: Qwen3MLPConfig) -> dict:
212
+ k1, k2, k3 = jax.random.split(key, 3)
213
+ def kaiming(k, shape):
214
+ return jax.random.normal(k, shape, dtype=jnp.float32) * jnp.sqrt(2.0 / shape[0])
215
+ return {
216
+ "w_gate": kaiming(k1, (cfg.in_dim, cfg.hidden_dim)),
217
+ "w_up": kaiming(k2, (cfg.in_dim, cfg.hidden_dim)),
218
+ "w_down": kaiming(k3, (cfg.hidden_dim, cfg.in_dim))
219
+ }
220
+
221
+ #===================================================================================
222
+ #Qwen3MoEBlock
223
+
224
+ def qwen3_moe_block(x: jax.Array, w: dict, cfg: Qwen3MoEBlockConfig) -> jax.Array:
225
+ """
226
+ Args:
227
+ x: [B, S, D_model]
228
+ Returns:
229
+ out: [B, S, D_model]
230
+ """
231
+ tokens = x.reshape(-1, x.shape[-1]) # (T, d)
232
+ router_logits = tokens @ w["router"] # (T, E)
233
+ probs = jax.nn.softmax(router_logits.astype(jnp.float32), -1)
234
+ topk_probs, topk_idx = jax.lax.top_k(probs, cfg.top_k) # (T, k)
235
+ topk_probs = topk_probs / topk_probs.sum(-1, keepdims=True) # norm_topk_prob
236
+ # scatter the renormalized probs back to (T, E), zeros elsewhere
237
+ weights = jnp.zeros_like(probs).at[
238
+ jnp.arange(tokens.shape[0])[:, None], topk_idx
239
+ ].set(topk_probs)
240
+ expert_out = jax.vmap(
241
+ lambda we: qwen3_mlp(tokens, we, cfg.mlp_conf)
242
+ )(w["experts"]) # (E, T, d)
243
+ out = jnp.einsum("te,etd->td", weights.astype(x.dtype), expert_out)
244
+ return out.reshape(x.shape)
245
+
246
+ def qwen3_moe_block_init(key, cfg: Qwen3MoEBlockConfig) -> dict:
247
+ k_router, k_experts = jax.random.split(key, 2)
248
+ expert_keys = jax.random.split(k_experts, cfg.num_experts)
249
+ return {
250
+ "router": jax.random.normal(k_router, (cfg.d_model, cfg.num_experts), dtype=jnp.bfloat16) * (cfg.d_model ** -0.5),
251
+ "experts": jax.vmap(lambda k: qwen3_mlp_init(k, cfg.mlp_conf))(expert_keys),
252
+ }
253
+
254
+ #===================================================================================
255
+ #Qwen3MoELayer
256
+
257
+ def qwen3_moe_layer(x: jax.Array, w: dict, cfg: Qwen3MoELayerConfig) -> jax.Array:
258
+ """
259
+ Args:
260
+ x: [B, S, D_model]
261
+ Returns:
262
+ out: [B, S, D_model]
263
+ """
264
+ residual = x
265
+ x = rms_norm(x, w["input_layernorm"], cfg.input_rms_conf)
266
+ x = qwen3_attn(x, w["attn"], cfg.attn_conf)
267
+ x = x + residual
268
+ residual = x
269
+ x = rms_norm(x, w["post_attn_layernorm"], cfg.post_attn_rms_conf)
270
+ x = qwen3_moe_block(x, w["mlp"], cfg.moe_block_conf)
271
+ x = residual + x
272
+ return x
273
+
274
+ def qwen3_moe_layer_init(key, cfg: Qwen3MoELayerConfig) -> dict:
275
+ attn_key, mlp_key, k3, k4 = jax.random.split(key, 4)
276
+ return {
277
+ "attn": qwen3_attn_init(attn_key, cfg.attn_conf),
278
+ "mlp": qwen3_moe_block_init(mlp_key, cfg.moe_block_conf),
279
+ "input_layernorm": rms_norm_init(k3, cfg.input_rms_conf),
280
+ "post_attn_layernorm": rms_norm_init(k4, cfg.post_attn_rms_conf)
281
+ }
282
+
283
+ #===================================================================================
284
+ #Qwen3MoeModel
285
+
286
+ def qwen3_moe_model(x: jax.Array, w: dict, cfg: Qwen3MoEModelConfig) -> jax.Array:
287
+ x = w["embeds"][x]
288
+ for l in w["decoder_blocks"]:
289
+ x = qwen3_moe_layer(
290
+ x,
291
+ l,
292
+ cfg.moe_layer_conf
293
+ )
294
+ x = rms_norm(x, w["norm"], cfg.final_rms_conf)
295
+ logits = x @ w["lm_head"]
296
+ return logits
297
+
298
+ def qwen3_moe_model_init(key, cfg: Qwen3MoEModelConfig) -> dict:
299
+ k_embed, k_layers, k_norm, k_head = jax.random.split(key, 4)
300
+ embeds = jax.random.normal(k_embed, shape=(cfg.vocab_size, cfg.d_model)) * (
301
+ cfg.d_model**-0.5
302
+ )
303
+ layer_keys = jax.random.split(k_layers, cfg.num_decoder_blocks)
304
+ decoder_blocks = [
305
+ qwen3_moe_layer_init(k, cfg.moe_layer_conf) for k in layer_keys
306
+ ]
307
+ norm = rms_norm_init(k_norm, cfg.final_rms_conf)
308
+ if cfg.tie_word_embeddings:
309
+ lm_head = embeds.T
310
+ else:
311
+ lm_head = jax.random.normal(k_head, shape=(cfg.d_model, cfg.vocab_size)) * (
312
+ cfg.d_model**-0.5
313
+ )
314
+
315
+ return {
316
+ "embeds": embeds,
317
+ "decoder_blocks": decoder_blocks,
318
+ "norm": norm,
319
+ "lm_head": lm_head,
320
+ }
321
+
322
+ #===================================================================================
adjaxt/model_maps.py ADDED
@@ -0,0 +1,27 @@
1
+ from adjaxt.sharding import *
2
+
3
+ QWEN3_MOE_WEIGHT_MAP = ModelWeightMap(
4
+ specs=[
5
+ WeightSpec("embeds", "model.embed_tokens.weight"),
6
+ WeightSpec("norm", "model.norm.weight"),
7
+ WeightSpec("lm_head", "lm_head.weight", transpose=True),
8
+
9
+ # Layer Norms
10
+ WeightSpec("decoder_blocks.{i}.input_layernorm", "model.layers.{i}.input_layernorm.weight"),
11
+ WeightSpec("decoder_blocks.{i}.post_attn_layernorm", "model.layers.{i}.post_attention_layernorm.weight"),
12
+
13
+ # Attention Projections & QK Norms
14
+ WeightSpec("decoder_blocks.{i}.attn.q_proj", "model.layers.{i}.self_attn.q_proj.weight", transpose=True),
15
+ WeightSpec("decoder_blocks.{i}.attn.k_proj", "model.layers.{i}.self_attn.k_proj.weight", transpose=True),
16
+ WeightSpec("decoder_blocks.{i}.attn.v_proj", "model.layers.{i}.self_attn.v_proj.weight", transpose=True),
17
+ WeightSpec("decoder_blocks.{i}.attn.o_proj", "model.layers.{i}.self_attn.o_proj.weight", transpose=True),
18
+ WeightSpec("decoder_blocks.{i}.attn.q_norm", "model.layers.{i}.self_attn.q_norm.weight"),
19
+ WeightSpec("decoder_blocks.{i}.attn.k_norm", "model.layers.{i}.self_attn.k_norm.weight"),
20
+
21
+ # MoE Router & Stacked Experts
22
+ WeightSpec("decoder_blocks.{i}.mlp.router", "model.layers.{i}.mlp.gate.weight", transpose=True),
23
+ WeightSpec("decoder_blocks.{i}.mlp.experts.w_gate", "model.layers.{i}.mlp.experts.{e}.gate_proj.weight", transpose=True),
24
+ WeightSpec("decoder_blocks.{i}.mlp.experts.w_up", "model.layers.{i}.mlp.experts.{e}.up_proj.weight", transpose=True),
25
+ WeightSpec("decoder_blocks.{i}.mlp.experts.w_down", "model.layers.{i}.mlp.experts.{e}.down_proj.weight", transpose=True),
26
+ ]
27
+ )
adjaxt/models.py ADDED
@@ -0,0 +1,148 @@
1
+ from adjaxt.config import *
2
+ from adjaxt.layers import *
3
+ from adjaxt.sharding import *
4
+ from adjaxt.model_maps import *
5
+ from huggingface_hub import snapshot_download, hf_hub_download
6
+ from typing import Callable
7
+
8
+ #=====================================================================================================
9
+ #Qwen3 MoE
10
+
11
+ def qwen3_moe_model_load(cfg: Qwen3MoEModelConfig, ckpt_path: str, token: Optional[str] = None):
12
+ """
13
+ Loads weights into a Qwen3 MoE JAX model from either a local path
14
+ or directly from a Hugging Face repository repo_id.
15
+ """
16
+ # 1. Resolve checkpoint path (download from HF hub if not a local folder/file)
17
+ if not os.path.exists(ckpt_path):
18
+ resolved_path = snapshot_download(
19
+ repo_id=ckpt_path,
20
+ allow_patterns=["*.safetensors", "*.json"],
21
+ token=token,
22
+ )
23
+ else:
24
+ resolved_path = ckpt_path
25
+
26
+ # 2. Extract loop/stack dimension limits dynamically from cfg
27
+ dim_sizes = {
28
+ "i": cfg.num_decoder_blocks,
29
+ "e": cfg.moe_layer_conf.moe_block_conf.num_experts,
30
+ }
31
+
32
+ # 3. Load flat safetensors and reconstruct the nested JAX tree
33
+ weights = load_checkpoint(
34
+ checkpoint_path=resolved_path,
35
+ weight_map=QWEN3_MOE_WEIGHT_MAP,
36
+ dim_sizes=dim_sizes,
37
+ )
38
+
39
+ # 4. Handle tied word embeddings if configured
40
+ if cfg.tie_word_embeddings:
41
+ weights["lm_head"] = weights["embeds"].T
42
+
43
+ return weights
44
+
45
+ def _get_act_fn(act_name: str) -> Callable[[jax.Array], jax.Array]:
46
+ act_map = {
47
+ "silu": jax.nn.silu,
48
+ "swish": jax.nn.swish,
49
+ "gelu": jax.nn.gelu,
50
+ "relu": jax.nn.relu,
51
+ }
52
+ if act_name not in act_map:
53
+ raise ValueError(f"Unsupported activation function: {act_name}")
54
+ return act_map[act_name]
55
+
56
+ def create_qwen3_moe_config(
57
+ hf_config: dict,
58
+ attn_implementation: StandardAttnImplementation = StandardAttnImplementation.XLA,
59
+ ) -> Qwen3MoEModelConfig:
60
+ """Builds an MoE Qwen3 configuration dataclass from a parsed HF config.json dictionary."""
61
+ hidden_size = hf_config["hidden_size"]
62
+ num_heads = hf_config["num_attention_heads"]
63
+ num_kv_heads = hf_config.get("num_key_value_heads", num_heads)
64
+ head_dim = hf_config.get("head_dim", hidden_size // num_heads)
65
+ rms_eps = hf_config.get("rms_norm_eps", 1e-6)
66
+ moe_intermediate_size = hf_config.get("moe_intermediate_size", hf_config.get("intermediate_size"))
67
+ num_layers = hf_config["num_hidden_layers"]
68
+ vocab_size = hf_config["vocab_size"]
69
+ rope_theta = float(hf_config.get("rope_theta", 1000000.0))
70
+ max_pos = hf_config.get("max_position_embeddings", 32768)
71
+ num_experts = hf_config.get("num_experts", hf_config.get("num_local_experts", 64))
72
+ top_k = hf_config.get("num_experts_per_tok", 8)
73
+ act_fn = _get_act_fn(hf_config.get("hidden_act", "silu"))
74
+
75
+ gqa_conf = GQAAttnConfig(
76
+ implementation=attn_implementation,
77
+ num_kv_groups=num_heads // num_kv_heads,
78
+ is_causal=True,
79
+ )
80
+ q_rms_conf = RMSNormConfig(dim=head_dim, eps=rms_eps)
81
+ k_rms_conf = RMSNormConfig(dim=head_dim, eps=rms_eps)
82
+
83
+ attn_conf = Qwen3AttnConfig(
84
+ gqa_conf=gqa_conf,
85
+ q_rms_conf=q_rms_conf,
86
+ k_rms_conf=k_rms_conf,
87
+ d=hidden_size,
88
+ head_dim=head_dim,
89
+ rope_theta=rope_theta,
90
+ max_position_embeddings=max_pos,
91
+ n_layers=num_layers,
92
+ )
93
+
94
+ mlp_conf = Qwen3MLPConfig(
95
+ act_fn=act_fn,
96
+ in_dim=hidden_size,
97
+ hidden_dim=moe_intermediate_size,
98
+ )
99
+
100
+ moe_block_conf = Qwen3MoEBlockConfig(
101
+ mlp_conf=mlp_conf,
102
+ top_k=top_k,
103
+ num_experts=num_experts,
104
+ d_model=hidden_size,
105
+ )
106
+
107
+ moe_layer_conf = Qwen3MoELayerConfig(
108
+ input_rms_conf=RMSNormConfig(dim=hidden_size, eps=rms_eps),
109
+ attn_conf=attn_conf,
110
+ post_attn_rms_conf=RMSNormConfig(dim=hidden_size, eps=rms_eps),
111
+ moe_block_conf=moe_block_conf,
112
+ )
113
+
114
+ return Qwen3MoEModelConfig(
115
+ moe_layer_conf=moe_layer_conf,
116
+ final_rms_conf=RMSNormConfig(dim=hidden_size, eps=rms_eps),
117
+ num_decoder_blocks=num_layers,
118
+ vocab_size=vocab_size,
119
+ d_model=hidden_size,
120
+ tie_word_embeddings=hf_config.get("tie_word_embeddings", False),
121
+ )
122
+
123
+
124
+ def qwen3_moe_config_from_pretrained(
125
+ model_name_or_path: str,
126
+ attn_implementation: StandardAttnImplementation = StandardAttnImplementation.XLA,
127
+ token: Optional[str] = None,
128
+ ) -> Qwen3MoEModelConfig:
129
+ """Loads config.json from local disk or HuggingFace Hub and constructs the corresponding dataclass."""
130
+ if os.path.exists(model_name_or_path):
131
+ config_file = (
132
+ model_name_or_path
133
+ if os.path.isfile(model_name_or_path)
134
+ else os.path.join(model_name_or_path, "config.json")
135
+ )
136
+ else:
137
+ config_file = hf_hub_download(
138
+ repo_id=model_name_or_path,
139
+ filename="config.json",
140
+ token=token,
141
+ )
142
+
143
+ with open(config_file, "r", encoding="utf-8") as f:
144
+ hf_config = json.load(f)
145
+
146
+ return create_qwen3_moe_config(hf_config, attn_implementation)
147
+
148
+ #=====================================================================================================
adjaxt/sharding.py ADDED
@@ -0,0 +1,251 @@
1
+ import functools
2
+ import glob
3
+ import itertools
4
+ import json
5
+ import operator
6
+ import os
7
+ import re
8
+ from dataclasses import dataclass
9
+ from typing import Any, Dict, List, Optional, Union
10
+
11
+ import jax
12
+ import jax.numpy as jnp
13
+ import numpy as np
14
+ from safetensors.numpy import load_file, save_file
15
+
16
+ # ==============================================================================
17
+ # 1. Spec & Model Definitions
18
+ # ==============================================================================
19
+
20
+ @dataclass
21
+ class WeightSpec:
22
+ jax_path: str
23
+ hf_path: str
24
+ transpose: bool = False
25
+
26
+ @dataclass
27
+ class ModelWeightMap:
28
+ specs: List[WeightSpec]
29
+
30
+ # ==============================================================================
31
+ # 2. Path Traversal & Pattern Parsing
32
+ # ==============================================================================
33
+
34
+ def get_nested_val(d: dict, path: List[Union[str, int]]) -> Any:
35
+ return functools.reduce(operator.getitem, path, d)
36
+
37
+ def set_nested_val(d: dict, path: List[Union[str, int]], val: Any) -> None:
38
+ def traverse_or_create(curr: Any, pair: tuple) -> Any:
39
+ key, next_key = pair
40
+ default_next = [] if isinstance(next_key, int) else {}
41
+
42
+ if isinstance(curr, list):
43
+ while len(curr) <= key:
44
+ curr.append(default_next.copy())
45
+ return curr[key]
46
+ return curr.setdefault(key, default_next)
47
+
48
+ lookahead_pairs = list(zip(path[:-1], path[1:]))
49
+ target_container = functools.reduce(traverse_or_create, lookahead_pairs, d)
50
+
51
+ last_key = path[-1]
52
+ if isinstance(target_container, list):
53
+ while len(target_container) <= last_key:
54
+ target_container.append(None)
55
+ target_container[last_key] = val
56
+
57
+ def parse_path_pattern(pattern: str, **kwargs) -> List[Union[str, int]]:
58
+ formatted = pattern.format(**kwargs)
59
+ return [int(t) if t.isdigit() else t for t in formatted.split(".")]
60
+
61
+ # ==============================================================================
62
+ # 3. Automatic Stack/Dimension Expansion
63
+ # ==============================================================================
64
+
65
+ def _expand_spec(spec: WeightSpec, dim_sizes: Dict[str, int]):
66
+ """
67
+ Finds placeholders shared in jax_path (loop dimensions) vs placeholders
68
+ only present in hf_path (dimensions that need stacking/unstacking).
69
+ """
70
+ jax_placeholders = set(re.findall(r"\{([a-zA-Z0-9_]+)\}", spec.jax_path))
71
+ hf_placeholders = set(re.findall(r"\{([a-zA-Z0-9_]+)\}", spec.hf_path))
72
+
73
+ shared_placeholders = list(jax_placeholders & hf_placeholders)
74
+ stack_placeholders = list(hf_placeholders - jax_placeholders)
75
+
76
+ # Validate that all required dimensions exist in dim_sizes
77
+ for p in shared_placeholders + stack_placeholders:
78
+ if p not in dim_sizes:
79
+ raise KeyError(f"Missing dimension limit for '{{{p}}}' in dim_sizes.")
80
+
81
+ # 1. Expand standard outer loop/layer combinations
82
+ if not shared_placeholders:
83
+ outer_combos = [{}]
84
+ else:
85
+ shared_ranges = [range(dim_sizes[p]) for p in shared_placeholders]
86
+ outer_combos = [dict(zip(shared_placeholders, idxs)) for idxs in itertools.product(*shared_ranges)]
87
+
88
+ # 2. Generate sub-combos for stacked dimensions (e.g., experts)
89
+ if not stack_placeholders:
90
+ stack_combos = [{}]
91
+ else:
92
+ stack_ranges = [range(dim_sizes[p]) for p in stack_placeholders]
93
+ stack_combos = [dict(zip(stack_placeholders, idxs)) for idxs in itertools.product(*stack_ranges)]
94
+
95
+ for outer in outer_combos:
96
+ yield spec, outer, stack_combos
97
+
98
+ # ==============================================================================
99
+ # 4. Checkpoint Loading & Saving
100
+ # ==============================================================================
101
+
102
+ def load_checkpoint(
103
+ checkpoint_path: str,
104
+ weight_map: ModelWeightMap,
105
+ dim_sizes: Optional[Dict[str, int]] = None
106
+ ) -> dict:
107
+ dim_sizes = dim_sizes or {}
108
+
109
+ if os.path.isdir(checkpoint_path):
110
+ index_file = os.path.join(checkpoint_path, "model.safetensors.index.json")
111
+ if os.path.exists(index_file):
112
+ with open(index_file, "r") as f:
113
+ index = json.load(f)
114
+ files = {os.path.join(checkpoint_path, v) for v in index["weight_map"].values()}
115
+ else:
116
+ files = glob.glob(os.path.join(checkpoint_path, "*.safetensors"))
117
+ else:
118
+ files = [checkpoint_path]
119
+
120
+ flat_tensors = {}
121
+ for f in files:
122
+ flat_tensors.update(load_file(f))
123
+
124
+ weights = {}
125
+ for spec in weight_map.specs:
126
+ for expanded_spec, outer_kwargs, stack_kwargs_list in _expand_spec(spec, dim_sizes):
127
+ jax_path_tokens = parse_path_pattern(expanded_spec.jax_path, **outer_kwargs)
128
+
129
+ # Case A: Multiple HuggingFace keys collapse/stack into one JAX tensor
130
+ if len(stack_kwargs_list) > 1:
131
+ stacked_tensors = []
132
+ for stack_kw in stack_kwargs_list:
133
+ hf_key = expanded_spec.hf_path.format(**outer_kwargs, **stack_kw)
134
+ t = flat_tensors[hf_key]
135
+ if expanded_spec.transpose:
136
+ t = t.T
137
+ stacked_tensors.append(t)
138
+ set_nested_val(weights, jax_path_tokens, jnp.stack(stacked_tensors, axis=0))
139
+
140
+ # Case B: Standard 1-to-1 weight mapping
141
+ else:
142
+ hf_key = expanded_spec.hf_path.format(**outer_kwargs)
143
+ t = flat_tensors[hf_key]
144
+ if expanded_spec.transpose:
145
+ t = t.T
146
+ set_nested_val(weights, jax_path_tokens, jnp.array(t))
147
+
148
+ return weights
149
+
150
+
151
+ def _parse_size_in_bytes(size: Union[int, str]) -> int:
152
+ """Converts human-readable size strings (e.g. '5GB', '500MB') to byte count."""
153
+ if isinstance(size, int):
154
+ return size
155
+ units = {"B": 1, "KB": 10**3, "MB": 10**6, "GB": 10**9, "TB": 10**12,
156
+ "KIB": 1 << 10, "MIB": 1 << 20, "GIB": 1 << 30, "TIB": 1 << 40}
157
+ match = re.match(r"^([0-9.]+)\s*([a-zA-Z]+)$", size.strip())
158
+ if not match:
159
+ raise ValueError(f"Invalid max_shard_size format: {size}")
160
+ val, unit = match.groups()
161
+ unit = unit.upper()
162
+ if unit not in units:
163
+ raise ValueError(f"Unknown size unit: {unit}")
164
+ return int(float(val) * units[unit])
165
+
166
+
167
+ def save_checkpoint(
168
+ weights: dict,
169
+ save_directory: str,
170
+ weight_map: ModelWeightMap,
171
+ dim_sizes: Optional[Dict[str, int]] = None,
172
+ max_shard_size: Union[str, int] = "5GB"
173
+ ) -> None:
174
+ """
175
+ Saves a nested JAX weight dictionary to single or multi-shard safetensors
176
+ with a HuggingFace-compatible index file.
177
+ """
178
+ dim_sizes = dim_sizes or {}
179
+ max_shard_bytes = _parse_size_in_bytes(max_shard_size)
180
+ os.makedirs(save_directory, exist_ok=True)
181
+
182
+ # 1. Flatten and convert JAX weights into numpy dict
183
+ flat_tensors: Dict[str, np.ndarray] = {}
184
+ for spec in weight_map.specs:
185
+ for expanded_spec, outer_kwargs, stack_kwargs_list in _expand_spec(spec, dim_sizes):
186
+ jax_path_tokens = parse_path_pattern(expanded_spec.jax_path, **outer_kwargs)
187
+ val = get_nested_val(weights, jax_path_tokens)
188
+
189
+ # Case A: Tensor unstacks across missing dimensions (e.g., experts)
190
+ if len(stack_kwargs_list) > 1:
191
+ for idx, stack_kw in enumerate(stack_kwargs_list):
192
+ hf_key = expanded_spec.hf_path.format(**outer_kwargs, **stack_kw)
193
+ t = val[idx]
194
+ if expanded_spec.transpose:
195
+ t = t.T
196
+ flat_tensors[hf_key] = np.asarray(jax.device_get(t))
197
+
198
+ # Case B: Standard 1-to-1 weight mapping
199
+ else:
200
+ hf_key = expanded_spec.hf_path.format(**outer_kwargs)
201
+ t = val
202
+ if expanded_spec.transpose:
203
+ t = t.T
204
+ flat_tensors[hf_key] = np.asarray(jax.device_get(t))
205
+
206
+ # 2. Partition tensors into shards based on max byte size
207
+ total_size = sum(t.nbytes for t in flat_tensors.values())
208
+
209
+ # Save as a single un-sharded file if within size limit
210
+ if total_size <= max_shard_bytes:
211
+ single_path = os.path.join(save_directory, "model.safetensors")
212
+ save_file(flat_tensors, single_path)
213
+ return
214
+
215
+ shards: List[Dict[str, np.ndarray]] = []
216
+ current_shard: Dict[str, np.ndarray] = {}
217
+ current_shard_size = 0
218
+
219
+ for name, tensor in flat_tensors.items():
220
+ tensor_bytes = tensor.nbytes
221
+ if current_shard and (current_shard_size + tensor_bytes > max_shard_bytes):
222
+ shards.append(current_shard)
223
+ current_shard = {}
224
+ current_shard_size = 0
225
+
226
+ current_shard[name] = tensor
227
+ current_shard_size += tensor_bytes
228
+
229
+ if current_shard:
230
+ shards.append(current_shard)
231
+
232
+ # 3. Write individual shard files and compile the index
233
+ num_shards = len(shards)
234
+ weight_to_file = {}
235
+
236
+ for idx, shard in enumerate(shards, start=1):
237
+ filename = f"model-{idx:05d}-of-{num_shards:05d}.safetensors"
238
+ file_path = os.path.join(save_directory, filename)
239
+ save_file(shard, file_path)
240
+
241
+ for tensor_name in shard.keys():
242
+ weight_to_file[tensor_name] = filename
243
+
244
+ # 4. Write model.safetensors.index.json
245
+ index_data = {
246
+ "metadata": {"total_size": total_size},
247
+ "weight_map": weight_to_file
248
+ }
249
+ index_path = os.path.join(save_directory, "model.safetensors.index.json")
250
+ with open(index_path, "w", encoding="utf-8") as f:
251
+ json.dump(index_data, f, indent=2)
adjaxt/train.py ADDED
File without changes
@@ -0,0 +1,22 @@
1
+ Metadata-Version: 2.5
2
+ Name: adjaxt
3
+ Version: 0.0.1
4
+ Summary: Distributed training with JAX on free compute, for humans
5
+ License-Expression: MIT
6
+ Requires-Python: >=3.11
7
+ Requires-Dist: datasets>=2.19
8
+ Requires-Dist: huggingface-hub>=0.23
9
+ Requires-Dist: jax>=0.4.30
10
+ Requires-Dist: numpy
11
+ Requires-Dist: optax>=0.2.2
12
+ Requires-Dist: orbax-checkpoint>=0.5
13
+ Requires-Dist: safetensors>=0.4
14
+ Provides-Extra: cuda
15
+ Requires-Dist: jax[cuda12]>=0.4.30; extra == 'cuda'
16
+ Provides-Extra: test
17
+ Requires-Dist: pytest; extra == 'test'
18
+ Requires-Dist: torch; extra == 'test'
19
+ Requires-Dist: transformers; extra == 'test'
20
+ Description-Content-Type: text/markdown
21
+
22
+ Hi, this package is in development, not ready yet
@@ -0,0 +1,15 @@
1
+ adjaxt/__init__.py,sha256=9ku486EWxqPGr1CgD7DGLwgnyAMM1bPAG00vonjLBsg,809
2
+ adjaxt/approx.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
3
+ adjaxt/checkpoint.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
4
+ adjaxt/config.py,sha256=C-F3N866PiSxOAWQw5N1vta49-ReH7YEt6_svSK17xM,2375
5
+ adjaxt/data.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
6
+ adjaxt/hub.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
7
+ adjaxt/init.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
8
+ adjaxt/layers.py,sha256=_B5nTfPgqDVYAAI5dJ_nZyD4R02sOQM9OrqGoU76QUI,10847
9
+ adjaxt/model_maps.py,sha256=U4xLXywDya1vaDkycE7flKsVJsDNCWxi_cBU2_ewAoE,1761
10
+ adjaxt/models.py,sha256=_q80rhx5sxNsKGLODl24NGkintjGIN4M27cucbb-9dI,5372
11
+ adjaxt/sharding.py,sha256=9hpNiQNZOpgcD8teNV1WI9R8pIWIfhIyNCtVJKGUm0w,9840
12
+ adjaxt/train.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
13
+ adjaxt-0.0.1.dist-info/METADATA,sha256=8A2_A-CD5Exbdf_iNa6EnfGYo_X30xihNiGL_S9jnRM,683
14
+ adjaxt-0.0.1.dist-info/WHEEL,sha256=zOwg4jB6zX2kU910N-cMawjivD6tO8NEWvE12je1bVk,87
15
+ adjaxt-0.0.1.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.0
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any