Skip to content
Open
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 README.md
Original file line number Diff line number Diff line change
Expand Up @@ -14,14 +14,14 @@ burn inference and training of (baby) [dragon hatchling](https://arxiv.org/abs/2

## features

- [x] mixture-of-expert routing
- [x] training benchmarks and reporting
- [ ] adaptive tool discovery
- [ ] conditional (deep) gating
- [ ] document-coherent dataloading and scale mixup
- [ ] episodic memory
- [ ] fused kernels
- [ ] hierarchical, memory-aware recurrent state
- [ ] mixture-of-expert routing
- [ ] multi-modal architecture
- [ ] neuromorphic backend
- [ ] rl reasoning training
Expand Down
4 changes: 2 additions & 2 deletions bench/inference.rs
Original file line number Diff line number Diff line change
Expand Up @@ -142,11 +142,11 @@ fn log_theoretical_profile(config: &BDHConfig, cfg: &InferenceConfig) {
}

fn compute_latent_per_head(config: &BDHConfig) -> usize {
(config.mlp_internal_dim_multiplier * config.n_embd) / config.n_head
config.latent_per_head()
}

fn compute_latent_total(config: &BDHConfig) -> usize {
compute_latent_per_head(config) * config.n_head
config.latent_total()
}

fn estimated_query_tensor_bytes(config: &BDHConfig, cfg: &InferenceConfig) -> u128 {
Expand Down
4 changes: 2 additions & 2 deletions bench/train.rs
Original file line number Diff line number Diff line change
Expand Up @@ -148,11 +148,11 @@ fn log_theoretical_profile(config: &BDHConfig, cfg: &TrainConfig) {
}

fn compute_latent_per_head(config: &BDHConfig) -> usize {
(config.mlp_internal_dim_multiplier * config.n_embd) / config.n_head
config.latent_per_head()
}

fn compute_latent_total(config: &BDHConfig) -> usize {
compute_latent_per_head(config) * config.n_head
config.latent_total()
}

criterion_group!(benches, training_step_bench);
Expand Down
1 change: 1 addition & 0 deletions config/base.toml
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ n_layer = 6
n_embd = 256
n_head = 4
mlp_internal_dim_multiplier = 128
experts = 2
dropout = 0.1
fused_kernels = true
use_alibi = true
3 changes: 3 additions & 0 deletions src/bin/infer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,9 @@ fn build_model_config(overrides: &ModelOverrides) -> BDHConfig {
if let Some(n_head) = overrides.n_head {
model_config.n_head = n_head;
}
if let Some(experts) = overrides.experts {
model_config.n_expert = experts.max(1);
}
if let Some(multiplier) = overrides.mlp_internal_dim_multiplier {
model_config.mlp_internal_dim_multiplier = multiplier;
}
Expand Down
3 changes: 3 additions & 0 deletions src/config/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,7 @@ pub struct ModelOverrides {
pub n_layer: Option<usize>,
pub n_embd: Option<usize>,
pub n_head: Option<usize>,
pub experts: Option<usize>,
pub mlp_internal_dim_multiplier: Option<usize>,
pub dropout: Option<f64>,
pub fused_kernels: Option<bool>,
Expand Down Expand Up @@ -224,6 +225,7 @@ mod tests {
"n_layer = 6",
"n_embd = 256",
"n_head = 4",
"experts = 2",
"mlp_internal_dim_multiplier = 128",
"dropout = 0.1",
"fused_kernels = false",
Expand Down Expand Up @@ -278,6 +280,7 @@ mod tests {
assert_eq!(config.model.n_layer, Some(6));
assert_eq!(config.model.n_embd, Some(320));
assert_eq!(config.model.n_head, Some(4));
assert_eq!(config.model.experts, Some(2));
assert_eq!(config.model.mlp_internal_dim_multiplier, Some(128));
assert_eq!(config.model.dropout, Some(0.1));
assert_eq!(config.model.fused_kernels, Some(true));
Expand Down
13 changes: 13 additions & 0 deletions src/model/bdh.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ use crate::kernel::{BlockPattern1d, relu_lowrank};

use super::attention::Attention;
use super::config::{BDHConfig, FusedKernelConfig};
use super::router::Router;

const LAYER_NORM_EPS: f32 = 1e-5;

Expand All @@ -19,11 +20,13 @@ pub struct BDH<B: Backend> {
n_embd: usize,
n_head: usize,
mlp_internal_dim_multiplier: usize,
n_expert: usize,
vocab_size: usize,
kernel: FusedKernelConfig,
embed: Embedding<B>,
dropout: Dropout,
attention: Attention<B>,
router: Router<B>,
encoder: Param<Tensor<B, 3>>,
encoder_v: Param<Tensor<B, 3>>,
decoder: Param<Tensor<B, 2>>,
Expand All @@ -43,6 +46,13 @@ impl<B: Backend> BDH<B> {
device,
&config.fused_kernels,
);
let router = Router::new(
config.n_expert,
config.n_head,
config.n_embd,
latent_per_head,
device,
);

let weight_init = |shape: [usize; 2]| {
Tensor::<B, 2>::random(shape, TensorDistribution::Normal(0.0, 0.02), device)
Expand All @@ -68,11 +78,13 @@ impl<B: Backend> BDH<B> {
n_embd: config.n_embd,
n_head: config.n_head,
mlp_internal_dim_multiplier: config.mlp_internal_dim_multiplier,
n_expert: config.n_expert,
vocab_size: config.vocab_size,
kernel: config.fused_kernels,
embed,
dropout,
attention,
router,
encoder,
encoder_v,
decoder,
Expand Down Expand Up @@ -131,6 +143,7 @@ impl<B: Backend> BDH<B> {
activation::relu(y_latent)
};
let xy_sparse = x_sparse * y_sparse;
let xy_sparse = self.router.route(state.clone(), xy_sparse);
let xy_sparse = self.dropout.forward(xy_sparse);

let mixed = xy_sparse.swap_dims(1, 2);
Expand Down
1 change: 1 addition & 0 deletions src/model/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ mod attention;
mod bdh;
mod config;
mod loss;
mod router;

pub use bdh::BDH;
pub use config::{BDHConfig, FusedKernelConfig};
Expand Down
79 changes: 79 additions & 0 deletions src/model/router.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
use burn::module::{Module, Param};
use burn::tensor::backend::Backend;
use burn::tensor::{Distribution as TensorDistribution, Tensor, activation};

#[derive(Module, Debug)]
pub struct Router<B: Backend> {
experts: usize,
n_head: usize,
latent_per_head: usize,
latent_per_expert: usize,
weight: Param<Tensor<B, 3>>,
bias: Param<Tensor<B, 2>>,
}

impl<B: Backend> Router<B> {
pub fn new(
experts: usize,
n_head: usize,
n_embd: usize,
latent_per_head: usize,
device: &B::Device,
) -> Self {
assert!(experts >= 1, "router requires at least one expert");
assert!(
latent_per_head % experts == 0,
"latent size {latent_per_head} must be divisible by experts {experts}"
);

let weight = Tensor::<B, 3>::random(
[n_head, n_embd, experts],
TensorDistribution::Normal(0.0, 0.02),
device,
);
let bias = Tensor::<B, 2>::zeros([n_head, experts], device);

Self {
experts,
n_head,
latent_per_head,
latent_per_expert: latent_per_head / experts,
weight: Param::from_tensor(weight),
bias: Param::from_tensor(bias),
}
}

pub fn route(&self, gating_input: Tensor<B, 4>, activations: Tensor<B, 4>) -> Tensor<B, 4> {
if self.experts == 1 {
return activations;
}

let weight = self.weight.val().unsqueeze_dim::<4>(0);
let mut logits = gating_input.matmul(weight);
let bias = self.bias.val().reshape([1, self.n_head, 1, self.experts]);
logits = logits + bias;

let routing = activation::softmax(logits, 3);

let [batch, heads, time, latent] = activations.shape().dims();
debug_assert_eq!(heads, self.n_head);
debug_assert_eq!(latent, self.latent_per_head);

let routed = activations.reshape([
batch,
self.n_head,
time,
self.experts,
self.latent_per_expert,
]);

let routing = routing.unsqueeze_dim::<5>(4);
let weighted = routing * routed;

weighted.reshape([batch, self.n_head, time, self.latent_per_head])
}

pub fn experts(&self) -> usize {
self.experts
}
}
3 changes: 3 additions & 0 deletions src/train.rs
Original file line number Diff line number Diff line change
Expand Up @@ -468,6 +468,9 @@ fn build_model_config(overrides: &ModelOverrides) -> BDHConfig {
if let Some(n_head) = overrides.n_head {
model_config.n_head = n_head;
}
if let Some(experts) = overrides.experts {
model_config.n_expert = experts.max(1);
}
if let Some(multiplier) = overrides.mlp_internal_dim_multiplier {
model_config.mlp_internal_dim_multiplier = multiplier;
}
Expand Down
Loading