use rlx::ir::shape;
use rlx::ir::GraphExt;
use rlx::ops::MaskKind;
use rlx::prelude::*;
#[derive(Clone, Copy, Debug)]
pub struct EncoderSpec {
pub b: usize,
pub t: usize,
pub d: usize,
pub n_classes: usize,
pub num_layers: usize,
pub num_heads: usize,
pub ffn_dim: usize,
pub dw_kernel: usize,
pub aux: bool,
}
fn s3(a: usize, b: usize, c: usize) -> Shape {
Shape::new(&[a, b, c], DType::F32)
}
fn linear(g: &mut Graph, x: NodeId, w: NodeId, b: NodeId) -> NodeId {
let y = g.mm(x, w);
let out_d = g.shape(y).dims()[2].unwrap_static();
let b3 = g.reshape_(b, vec![1, 1, out_d as i64]);
g.add(y, b3)
}
pub fn build_conformer_tail_graph(spec: &EncoderSpec) -> Graph {
let mut g = Graph::new("brain2qwerty_conformer_tail");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let mut x = z_in;
for layer in 0..spec.num_layers {
x = conformer_layer(&mut g, x, layer, spec);
}
let z_final = x;
let mut c_in = z_final;
if spec.aux {
let ln_w = g.param(
"shared_layer_norm.weight",
Shape::new(&[spec.d], DType::F32),
);
let ln_b = g.param("shared_layer_norm.bias", Shape::new(&[spec.d], DType::F32));
c_in = g.ln(c_in, ln_w, ln_b, 1e-5);
}
let out_w = g.param(
"output_layer.weight",
Shape::new(&[spec.d, spec.n_classes], DType::F32),
);
let out_b = g.param(
"output_layer.bias",
Shape::new(&[spec.n_classes], DType::F32),
);
let c_out = linear(&mut g, c_in, out_w, out_b);
g.set_outputs(vec![z_final, c_out]);
g
}
pub fn build_self_attn_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_self_attn");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}");
let sa_ln_w = g.param(
format!("{p}.self_attn_layer_norm.weight"),
Shape::new(&[spec.d], DType::F32),
);
let sa_ln_b = g.param(
format!("{p}.self_attn_layer_norm.bias"),
Shape::new(&[spec.d], DType::F32),
);
let xn = g.ln(z_in, sa_ln_w, sa_ln_b, 1e-5);
let out = self_attn(&mut g, xn, layer, spec);
g.set_outputs(vec![out]);
g
}
pub fn build_conv_p1_glu_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_conv_p1_glu");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}.conv_module");
let d = spec.d;
let b = spec.b;
let t = spec.t;
let ln_w = g.param(
format!("{p}.layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let ln_b = g.param(format!("{p}.layer_norm.bias"), Shape::new(&[d], DType::F32));
let x = g.ln(z_in, ln_w, ln_b, 1e-5);
let x4 = bct_to_nchw(&mut g, x, b, d, t);
let p1_w = g.param(
format!("{p}.sequential.0.weight"),
Shape::new(&[2 * d, d, 1, 1], DType::F32),
);
let p1_b = g.param(
format!("{p}.sequential.0.bias"),
Shape::new(&[2 * d], DType::F32),
);
let y = g.conv2d(x4, p1_w, [1, 1], [1, 1], [0, 0], [1, 1], 1);
let y = add_bias_nchw(&mut g, y, p1_b, 2 * d);
let y0 = g.narrow_(y, 1, 0, d);
let y1 = g.narrow_(y, 1, d, d);
let gate = sigmoid(&mut g, y1);
let glu = g.mul(y0, gate);
let out = nchw_to_btd(&mut g, glu, b, d, t);
g.set_outputs(vec![out]);
g
}
pub fn build_conv_dw_gn_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_conv_dw_gn");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}.conv_module");
let d = spec.d;
let b = spec.b;
let t = spec.t;
let k = spec.dw_kernel;
let pad = k / 2;
let glu = bct_to_nchw(&mut g, z_in, b, d, t);
let dw_w = g.param(
format!("{p}.sequential.2.weight"),
Shape::new(&[d, 1, 1, k], DType::F32),
);
let dw_b = g.param(
format!("{p}.sequential.2.bias"),
Shape::new(&[d], DType::F32),
);
let mut dw = g.conv2d(glu, dw_w, [1, k], [1, 1], [0, pad], [1, 1], d);
dw = add_bias_nchw(&mut g, dw, dw_b, d);
let gn_w = g.param(
format!("{p}.sequential.3.weight"),
Shape::new(&[d], DType::F32),
);
let gn_b = g.param(
format!("{p}.sequential.3.bias"),
Shape::new(&[d], DType::F32),
);
let dw_btd = nchw_to_btd(&mut g, dw, b, d, t);
let out = group_norm_1group_btd(&mut g, dw_btd, gn_w, gn_b, d, 1e-5);
g.set_outputs(vec![out]);
g
}
pub fn build_conv_dw_only_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_conv_dw_only");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}.conv_module");
let d = spec.d;
let b = spec.b;
let t = spec.t;
let k = spec.dw_kernel;
let pad = k / 2;
let glu = bct_to_nchw(&mut g, z_in, b, d, t);
let dw_w = g.param(
format!("{p}.sequential.2.weight"),
Shape::new(&[d, 1, 1, k], DType::F32),
);
let dw_b = g.param(
format!("{p}.sequential.2.bias"),
Shape::new(&[d], DType::F32),
);
let mut dw = g.conv2d(glu, dw_w, [1, k], [1, 1], [0, pad], [1, 1], d);
dw = add_bias_nchw(&mut g, dw, dw_b, d);
let out = nchw_to_btd(&mut g, dw, b, d, t);
g.set_outputs(vec![out]);
g
}
pub fn build_conv_dw_tail_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_conv_dw_tail");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}.conv_module");
let d = spec.d;
let b = spec.b;
let t = spec.t;
let k = spec.dw_kernel;
let pad = k / 2;
let glu = bct_to_nchw(&mut g, z_in, b, d, t);
let dw_w = g.param(
format!("{p}.sequential.2.weight"),
Shape::new(&[d, 1, 1, k], DType::F32),
);
let dw_b = g.param(
format!("{p}.sequential.2.bias"),
Shape::new(&[d], DType::F32),
);
let mut dw = g.conv2d(glu, dw_w, [1, k], [1, 1], [0, pad], [1, 1], d);
dw = add_bias_nchw(&mut g, dw, dw_b, d);
let gn_w = g.param(
format!("{p}.sequential.3.weight"),
Shape::new(&[d], DType::F32),
);
let gn_b = g.param(
format!("{p}.sequential.3.bias"),
Shape::new(&[d], DType::F32),
);
let dw_btd = nchw_to_btd(&mut g, dw, b, d, t);
let dw_btd = group_norm_1group_btd(&mut g, dw_btd, gn_w, gn_b, d, 1e-5);
let dw_btd = g.silu(dw_btd);
let dw4 = bct_to_nchw(&mut g, dw_btd, b, d, t);
let p2_w = g.param(
format!("{p}.sequential.5.weight"),
Shape::new(&[d, d, 1, 1], DType::F32),
);
let p2_b = g.param(
format!("{p}.sequential.5.bias"),
Shape::new(&[d], DType::F32),
);
let out = g.conv2d(dw4, p2_w, [1, 1], [1, 1], [0, 0], [1, 1], 1);
let out = add_bias_nchw(&mut g, out, p2_b, d);
let out = nchw_to_btd(&mut g, out, b, d, t);
g.set_outputs(vec![out]);
g
}
pub fn build_conv_module_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_conv_module");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let out = conv_module(&mut g, z_in, layer, spec);
g.set_outputs(vec![out]);
g
}
pub fn build_pre_conv_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_pre_conv");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}");
let d = spec.d;
let half = scalar(&mut g, 0.5);
let residual = z_in;
let ffn1 = ffn_block(&mut g, z_in, &format!("{p}.ffn1"), spec.ffn_dim);
let scaled = g.mul(ffn1, half);
let x = g.add(scaled, residual);
let residual = x;
let sa_ln_w = g.param(
format!("{p}.self_attn_layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let sa_ln_b = g.param(
format!("{p}.self_attn_layer_norm.bias"),
Shape::new(&[d], DType::F32),
);
let xn = g.ln(x, sa_ln_w, sa_ln_b, 1e-5);
let attn = self_attn(&mut g, xn, layer, spec);
let out = g.add(attn, residual);
g.set_outputs(vec![out]);
g
}
pub fn build_post_conv_graph(spec: &EncoderSpec, layer: usize) -> Graph {
let mut g = Graph::new("brain2qwerty_post_conv");
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}");
let d = spec.d;
let half = scalar(&mut g, 0.5);
let residual = z_in;
let ffn2 = ffn_block(&mut g, z_in, &format!("{p}.ffn2"), spec.ffn_dim);
let scaled = g.mul(ffn2, half);
let x = g.add(scaled, residual);
let fln_w = g.param(
format!("{p}.final_layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let fln_b = g.param(
format!("{p}.final_layer_norm.bias"),
Shape::new(&[d], DType::F32),
);
let out = g.ln(x, fln_w, fln_b, 1e-5);
g.set_outputs(vec![out]);
g
}
pub fn build_ffn_block_graph(spec: &EncoderSpec, layer: usize, which: &str) -> Graph {
let mut g = Graph::new(format!("brain2qwerty_ffn_{which}"));
let z_in = g.input("z", s3(spec.b, spec.t, spec.d));
let p = format!("transformer.conformer_layers.{layer}.{which}");
let out = ffn_block(&mut g, z_in, &p, spec.ffn_dim);
g.set_outputs(vec![out]);
g
}
fn conformer_layer(g: &mut Graph, x: NodeId, layer: usize, spec: &EncoderSpec) -> NodeId {
let p = format!("transformer.conformer_layers.{layer}");
let d = spec.d;
let half = scalar(g, 0.5);
let residual = x;
let ffn1 = ffn_block(g, x, &format!("{p}.ffn1"), spec.ffn_dim);
let scaled = g.mul(ffn1, half);
let x = g.add(scaled, residual);
let residual = x;
let sa_ln_w = g.param(
format!("{p}.self_attn_layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let sa_ln_b = g.param(
format!("{p}.self_attn_layer_norm.bias"),
Shape::new(&[d], DType::F32),
);
let xn = g.ln(x, sa_ln_w, sa_ln_b, 1e-5);
let attn = self_attn(g, xn, layer, spec);
let x = g.add(attn, residual);
let residual = x;
let conv = conv_module(g, x, layer, spec);
let x = g.add(conv, residual);
let residual = x;
let ffn2 = ffn_block(g, x, &format!("{p}.ffn2"), spec.ffn_dim);
let scaled2 = g.mul(ffn2, half);
let x = g.add(scaled2, residual);
let fln_w = g.param(
format!("{p}.final_layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let fln_b = g.param(
format!("{p}.final_layer_norm.bias"),
Shape::new(&[d], DType::F32),
);
g.ln(x, fln_w, fln_b, 1e-5)
}
fn scalar(g: &mut Graph, v: f32) -> NodeId {
g.append_node(
Op::Constant {
data: v.to_le_bytes().to_vec(),
},
vec![],
Shape::new(&[1], DType::F32),
None,
)
}
fn group_norm_1group_btd(
g: &mut Graph,
x: NodeId,
weight: NodeId,
bias: NodeId,
d: usize,
eps: f32,
) -> NodeId {
let eps_n = scalar(g, eps);
let one = scalar(g, 1.0);
let mean = g.mean(x, vec![1, 2], true);
let centered = g.sub(x, mean);
let sq = g.mul(centered, centered);
let var = g.mean(sq, vec![1, 2], true);
let var_eps = g.add(var, eps_n);
let inv_denom = g.sqrt(var_eps);
let inv = g.div(one, inv_denom);
let normed = g.mul(centered, inv);
let w = g.reshape_(weight, vec![1, 1, d as i64]);
let b = g.reshape_(bias, vec![1, 1, d as i64]);
let scaled = g.mul(normed, w);
g.add(scaled, b)
}
fn ffn_block(g: &mut Graph, x: NodeId, prefix: &str, ffn_dim: usize) -> NodeId {
let dim = g.shape(x).dims()[2].unwrap_static();
let ln_w = g.param(
format!("{prefix}.sequential.0.weight"),
Shape::new(&[dim], DType::F32),
);
let ln_b = g.param(
format!("{prefix}.sequential.0.bias"),
Shape::new(&[dim], DType::F32),
);
let x = g.ln(x, ln_w, ln_b, 1e-5);
let w1 = g.param(
format!("{prefix}.sequential.1.weight"),
Shape::new(&[dim, ffn_dim], DType::F32),
);
let b1 = g.param(
format!("{prefix}.sequential.1.bias"),
Shape::new(&[ffn_dim], DType::F32),
);
let mm1 = linear(g, x, w1, b1);
let h = g.silu(mm1);
let w2 = g.param(
format!("{prefix}.sequential.4.weight"),
Shape::new(&[ffn_dim, dim], DType::F32),
);
let b2 = g.param(
format!("{prefix}.sequential.4.bias"),
Shape::new(&[dim], DType::F32),
);
linear(g, h, w2, b2)
}
fn self_attn(g: &mut Graph, x: NodeId, layer: usize, spec: &EncoderSpec) -> NodeId {
let p = format!("transformer.conformer_layers.{layer}.self_attn");
let d = spec.d;
let nh = spec.num_heads;
let dh = d / nh;
let in_w = g.param(
format!("{p}.in_proj.weight"),
Shape::new(&[d, 3 * d], DType::F32),
);
let in_b = g.param(
format!("{p}.in_proj.bias"),
Shape::new(&[3 * d], DType::F32),
);
let wo = g.param(
format!("{p}.out_proj.weight"),
Shape::new(&[d, d], DType::F32),
);
let bo = g.param(format!("{p}.out_proj.bias"), Shape::new(&[d], DType::F32));
let qkv = linear(g, x, in_w, in_b);
let q = g.narrow_(qkv, 2, 0, d);
let k = g.narrow_(qkv, 2, d, d);
let v = g.narrow_(qkv, 2, 2 * d, d);
let b = spec.b;
let s = spec.t;
let q4 = g.reshape_(q, vec![b as i64, s as i64, nh as i64, dh as i64]);
let k4 = g.reshape_(k, vec![b as i64, s as i64, nh as i64, dh as i64]);
let v4 = g.reshape_(v, vec![b as i64, s as i64, nh as i64, dh as i64]);
let q_bhs = g.transpose_(q4, vec![0, 2, 1, 3]);
let k_bhs = g.transpose_(k4, vec![0, 2, 1, 3]);
let v_bhs = g.transpose_(v4, vec![0, 2, 1, 3]);
let attn_shape = shape::attention_shape(g.shape(q_bhs));
let attn = g.attention_kind(q_bhs, k_bhs, v_bhs, nh, dh, MaskKind::None, attn_shape);
let attn_bsh = g.transpose_(attn, vec![0, 2, 1, 3]);
let attn_3 = g.reshape_(attn_bsh, vec![b as i64, s as i64, (nh * dh) as i64]);
linear(g, attn_3, wo, bo)
}
fn sigmoid(g: &mut Graph, x: NodeId) -> NodeId {
let neg_one = scalar(g, -1.0);
let neg = g.mul(x, neg_one);
let e = g.exp(neg);
let one = scalar(g, 1.0);
let denom = g.add(one, e);
g.div(one, denom)
}
fn bct_to_nchw(g: &mut Graph, x: NodeId, b: usize, d: usize, t: usize) -> NodeId {
let x_bdt = g.transpose_(x, vec![0, 2, 1]);
g.reshape_(x_bdt, vec![b as i64, d as i64, 1, t as i64])
}
fn nchw_to_btd(g: &mut Graph, x: NodeId, b: usize, d: usize, t: usize) -> NodeId {
let x3 = g.reshape_(x, vec![b as i64, d as i64, t as i64]);
g.transpose_(x3, vec![0, 2, 1])
}
fn add_bias_nchw(g: &mut Graph, x: NodeId, bias: NodeId, channels: usize) -> NodeId {
let b = g.reshape_(bias, vec![1, channels as i64, 1, 1]);
g.add(x, b)
}
fn conv_module(g: &mut Graph, x: NodeId, layer: usize, spec: &EncoderSpec) -> NodeId {
let _ = layer;
let p = format!("transformer.conformer_layers.{layer}.conv_module");
let d = spec.d;
let b = spec.b;
let t = spec.t;
let k = spec.dw_kernel;
let pad = k / 2;
let ln_w = g.param(
format!("{p}.layer_norm.weight"),
Shape::new(&[d], DType::F32),
);
let ln_b = g.param(format!("{p}.layer_norm.bias"), Shape::new(&[d], DType::F32));
let x = g.ln(x, ln_w, ln_b, 1e-5);
let x4 = bct_to_nchw(g, x, b, d, t);
let p1_w = g.param(
format!("{p}.sequential.0.weight"),
Shape::new(&[2 * d, d, 1, 1], DType::F32),
);
let p1_b = g.param(
format!("{p}.sequential.0.bias"),
Shape::new(&[2 * d], DType::F32),
);
let y = g.conv2d(x4, p1_w, [1, 1], [1, 1], [0, 0], [1, 1], 1);
let y = add_bias_nchw(g, y, p1_b, 2 * d);
let y0 = g.narrow_(y, 1, 0, d);
let y1 = g.narrow_(y, 1, d, d);
let gate = sigmoid(g, y1);
let glu = g.mul(y0, gate);
let dw_w = g.param(
format!("{p}.sequential.2.weight"),
Shape::new(&[d, 1, 1, k], DType::F32),
);
let dw_b = g.param(
format!("{p}.sequential.2.bias"),
Shape::new(&[d], DType::F32),
);
let mut dw = g.conv2d(glu, dw_w, [1, k], [1, 1], [0, pad], [1, 1], d);
dw = add_bias_nchw(g, dw, dw_b, d);
let gn_w = g.param(
format!("{p}.sequential.3.weight"),
Shape::new(&[d], DType::F32),
);
let gn_b = g.param(
format!("{p}.sequential.3.bias"),
Shape::new(&[d], DType::F32),
);
let dw_btd = nchw_to_btd(g, dw, b, d, t);
let dw_btd = group_norm_1group_btd(g, dw_btd, gn_w, gn_b, d, 1e-5);
let dw_btd = g.silu(dw_btd);
let dw = bct_to_nchw(g, dw_btd, b, d, t);
let p2_w = g.param(
format!("{p}.sequential.5.weight"),
Shape::new(&[d, d, 1, 1], DType::F32),
);
let p2_b = g.param(
format!("{p}.sequential.5.bias"),
Shape::new(&[d], DType::F32),
);
let out = g.conv2d(dw, p2_w, [1, 1], [1, 1], [0, 0], [1, 1], 1);
let out = add_bias_nchw(g, out, p2_b, d);
nchw_to_btd(g, out, b, d, t)
}