use anyhow::Result;
use rlx_core::config::BertConfig;
use rlx_core::weight_map::WeightMap;
use rlx_ir::Graph;
use std::collections::HashMap;
pub fn build_bert_graph_sized(
cfg: &BertConfig,
weights: &mut WeightMap,
batch: usize,
seq: usize,
) -> Result<(Graph, HashMap<String, Vec<f32>>)> {
rlx_core::flow_util::graph_from_built(crate::flow::build_bert_built(cfg, weights, batch, seq)?)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_tiny_bert_graph() {
let cfg = BertConfig {
vocab_size: 100,
hidden_size: 64,
num_hidden_layers: 1,
num_attention_heads: 2,
intermediate_size: 256,
max_position_embeddings: 32,
type_vocab_size: 2,
layer_norm_eps: 1e-12,
hidden_act: "gelu".into(),
};
let mut wm = crate::test_support::tiny_bert_weights(&cfg);
let (graph, params) = build_bert_graph_sized(&cfg, &mut wm, 1, 1).unwrap();
let errors = rlx_ir::verify::verify(&graph);
assert!(errors.is_empty(), "verification errors: {errors:?}");
assert!(
params.len() >= 15,
"expected 15+ params, got {}",
params.len()
);
assert!(!graph.outputs.is_empty());
}
}