import jax
import jax.numpy as jnp
from typing import NamedTuple, Any, Callable
class ModuleOutput(NamedTuple):
"""Explicit interface: modules only produce outputs of known type."""
value: jnp.ndarray
state: dict
class StrictModule:
"""Base pattern for strictly encapsulated modules."""
def __init__(self, config: dict):
"""Modules are configured at construction, not modified later."""
self.config = config
self.params = None
def initialize_parameters(self, key: jax.random.PRNGKey, input_shape: tuple):
"""Explicit parameter initialization with shape contract."""
raise NotImplementedError
def __call__(
self,
inputs: jnp.ndarray,
training: bool = True,
context: dict = None
) -> ModuleOutput:
"""
Single call signature regardless of module type.
No hidden dependencies on module type or inheritance.
"""
raise NotImplementedError
class LinearLayer(StrictModule):
"""Example: strictly encapsulated linear layer."""
def __init__(self, config: dict):
super().__init__(config)
self.output_dim = config.get("output_dim", 768)
self.use_bias = config.get("use_bias", True)
def initialize_parameters(self, key, input_shape):
"""Deterministic initialization based on input shape."""
input_dim = input_shape[-1]
key_w, key_b = jax.random.split(key)
self.params = {
"weight": jax.random.normal(key_w, (input_dim, self.output_dim)) * jnp.sqrt(2.0 / input_dim),
}
if self.use_bias:
self.params["bias"] = jnp.zeros(self.output_dim)
def __call__(self, inputs: jnp.ndarray, training: bool = True, context: dict = None) -> ModuleOutput:
"""Standard linear transformation with no side effects."""
output = jnp.dot(inputs, self.params["weight"])
if self.use_bias:
output = output + self.params["bias"]
return ModuleOutput(
value=output,
state={}
)
class AttentionLayer(StrictModule):
"""Example: strictly encapsulated attention."""
def __init__(self, config: dict):
super().__init__(config)
self.num_heads = config.get("num_heads", 12)
self.head_dim = config.get("hidden_dim", 768) // self.num_heads
def initialize_parameters(self, key, input_shape):
hidden_dim = input_shape[-1]
key_q, key_k, key_v, key_out = jax.random.split(key, 4)
scale = jnp.sqrt(2.0 / hidden_dim)
self.params = {
"q_proj": jax.random.normal(key_q, (hidden_dim, hidden_dim)) * scale,
"k_proj": jax.random.normal(key_k, (hidden_dim, hidden_dim)) * scale,
"v_proj": jax.random.normal(key_v, (hidden_dim, hidden_dim)) * scale,
"out_proj": jax.random.normal(key_out, (hidden_dim, hidden_dim)) * scale,
}
def __call__(self, inputs: jnp.ndarray, training: bool = True, context: dict = None) -> ModuleOutput:
batch, seq_len, hidden_dim = inputs.shape
Q = jnp.dot(inputs, self.params["q_proj"])
K = jnp.dot(inputs, self.params["k_proj"])
V = jnp.dot(inputs, self.params["v_proj"])
Q = Q.reshape(batch, seq_len, self.num_heads, self.head_dim).transpose(0, 2, 1, 3)
K = K.reshape(batch, seq_len, self.num_heads, self.head_dim).transpose(0, 2, 1, 3)
V = V.reshape(batch, seq_len, self.num_heads, self.head_dim).transpose(0, 2, 1, 3)
scores = jnp.matmul(Q, K.transpose(0, 1, 3, 2)) / jnp.sqrt(self.head_dim)
attn_weights = jax.nn.softmax(scores, axis=-1)
attn_output = jnp.matmul(attn_weights, V)
attn_output = attn_output.transpose(0, 2, 1, 3).reshape(batch, seq_len, hidden_dim)
output = jnp.dot(attn_output, self.params["out_proj"])
return ModuleOutput(
value=output,
state={}
)
class TransformerBlock(StrictModule):
"""Composing modules: attention + FFN block."""
def __init__(self, config: dict):
super().__init__(config)
self.attention = AttentionLayer(config)
self.ffn_layer1 = LinearLayer({"output_dim": config.get("ffn_dim", 3072)})
self.ffn_layer2 = LinearLayer({"output_dim": config.get("hidden_dim", 768)})
self.norm1 = config.get("norm_fn", "layer_norm")
def initialize_parameters(self, key, input_shape):
key_attn, key_ffn1, key_ffn2 = jax.random.split(key, 3)
self.attention.initialize_parameters(key_attn, input_shape)
self.ffn_layer1.initialize_parameters(key_ffn1, input_shape)
ffn_output_shape = input_shape[:-1] + (self.config.get("ffn_dim", 3072),)
self.ffn_layer2.initialize_parameters(key_ffn2, ffn_output_shape)
self.params = {
"attention": self.attention.params,
"ffn": {
"layer1": self.ffn_layer1.params,
"layer2": self.ffn_layer2.params,
}
}
def __call__(self, inputs: jnp.ndarray, training: bool = True, context: dict = None) -> ModuleOutput:
self.attention.params = self.params["attention"]
attn_out = self.attention(inputs, training, context)
x = inputs + attn_out.value
x_norm = jax.nn.layer_norm(x)
self.ffn_layer1.params = self.params["ffn"]["layer1"]
ffn_out1 = self.ffn_layer1(x_norm, training, context)
ffn_out1_activated = jax.nn.gelu(ffn_out1.value)
self.ffn_layer2.params = self.params["ffn"]["layer2"]
ffn_out2 = self.ffn_layer2(ffn_out1_activated, training, context)
x = x + ffn_out2.value
return ModuleOutput(
value=x,
state={**attn_out.state, **ffn_out2.state}
)
class TransformerModel(StrictModule):
"""Full transformer: composition of blocks."""
def __init__(self, config: dict):
super().__init__(config)
self.num_layers = config.get("num_layers", 12)
self.blocks = [TransformerBlock(config) for _ in range(self.num_layers)]
def initialize_parameters(self, key, input_shape):
keys = jax.random.split(key, self.num_layers)
self.params = {}
current_shape = input_shape
for i, block in enumerate(self.blocks):
block.initialize_parameters(keys[i], current_shape)
self.params[f"block_{i}"] = block.params
def __call__(self, inputs: jnp.ndarray, training: bool = True, context: dict = None) -> ModuleOutput:
x = inputs
all_states = {}
for i, block in enumerate(self.blocks):
block.params = self.params[f"block_{i}"]
output = block(x, training, context)
x = output.value
all_states.update(output.state)
return ModuleOutput(
value=x,
state=all_states
)
import optax
class HardwareAgnosticTrainer:
"""Training loop transparent to hardware backend."""
def __init__(self, model: StrictModule, config: dict):
self.model = model
self.config = config
self.optimizer = optax.adam(learning_rate=config.get("learning_rate", 1e-3))
@jax.jit
def train_step(self, params: dict, batch: dict, opt_state: dict):
"""Single training step (JIT-compiled for any hardware)."""
def loss_fn(params):
self.model.params = params
output = self.model(batch["input_ids"], training=True)
logits = output.value
targets = batch["labels"]
loss = optax.softmax_cross_entropy_with_integer_labels(logits, targets).mean()
return loss
loss, grads = jax.value_and_grad(loss_fn)(params)
updates, opt_state = self.optimizer.update(grads, opt_state)
params = optax.apply_updates(params, updates)
return params, opt_state, loss
def train_epoch(self, train_loader, num_epochs: int = 3):
"""Training epoch (runs on whatever hardware JAX detects)."""
key = jax.random.PRNGKey(0)
dummy_input = jnp.ones((1, 512, 768))
self.model.initialize_parameters(key, dummy_input.shape)
opt_state = self.optimizer.init(self.model.params)
for epoch in range(num_epochs):
for batch in train_loader:
self.model.params, opt_state, loss = self.train_step(
self.model.params, batch, opt_state
)
print(f"Epoch {epoch + 1}/{num_epochs} | Loss: {loss:.4f}")
class LoC_ComplexityAnalyzer:
"""Measure lines of code required to add features."""
@staticmethod
def analyze_feature_addition(feature_name: str, module_count: int) -> dict:
"""
Hypothetical feature addition: Add RotaryPositionalEmbedding to attention layers.
In traditional frameworks (subtyping): O(N) where N = module count
In AXLearn (strict encapsulation): O(1)
"""
axlearn_changes = {
"files_modified": 1,
"new_classes": 1,
"modified_methods": 1,
"total_loc_added": 50,
"complexity": "O(1)"
}
traditional_changes = {
"files_modified": module_count,
"modified_methods": module_count,
"total_loc_added": module_count * 10,
"complexity": f"O(N) where N={module_count}"
}
return {
"axlearn": axlearn_changes,
"traditional": traditional_changes,
"complexity_improvement": module_count
}