MCPcopy Create free account
hub / github.com/evilsocket/cake / load

Method load

cake-core/src/models/luxtts/zipformer_layer.rs:45–107  ·  view source on GitHub ↗
(
        dim: usize,
        ff_dim: usize,
        num_heads: usize,
        query_head_dim: usize,
        value_head_dim: usize,
        pos_dim: usize,
        pos_head_dim: usize,
        cnn_ke

Source from the content-addressed store, hash-verified

43impl ZipformerEncoderLayer {
44 #[allow(clippy::too_many_arguments)]
45 pub fn load(
46 dim: usize,
47 ff_dim: usize,
48 num_heads: usize,
49 query_head_dim: usize,
50 value_head_dim: usize,
51 pos_dim: usize,
52 pos_head_dim: usize,
53 cnn_kernel: usize,
54 vb: VarBuilder,
55 backend: Arc<dyn ComputeBackend>,
56 ) -> Result<Self> {
57 let norm = BiasNorm::load(dim, vb.pp("norm"))?;
58
59 // Three feed-forward modules with different intermediate sizes
60 let ff1_dim = ff_dim * 3 / 4;
61 let ff2_dim = ff_dim;
62 let ff3_dim = ff_dim * 5 / 4;
63 let feed_forward1 = FeedforwardModule::load(dim, ff1_dim, vb.pp("feed_forward1"), backend.clone())?;
64 let feed_forward2 = FeedforwardModule::load(dim, ff2_dim, vb.pp("feed_forward2"), backend.clone())?;
65 let feed_forward3 = FeedforwardModule::load(dim, ff3_dim, vb.pp("feed_forward3"), backend.clone())?;
66
67 // Attention weights (computes Q, K, pos -> attention matrix)
68 let self_attn_weights = RelPositionMultiheadAttentionWeights::load(
69 dim,
70 num_heads,
71 query_head_dim,
72 pos_head_dim,
73 pos_dim,
74 vb.pp("self_attn_weights"),
75 backend.clone(),
76 )?;
77
78 // Two self-attention modules (apply weights to values)
79 let self_attn1 = SelfAttention::load(dim, num_heads, value_head_dim, vb.pp("self_attn1"), backend.clone())?;
80 let self_attn2 = SelfAttention::load(dim, num_heads, value_head_dim, vb.pp("self_attn2"), backend.clone())?;
81
82 // Nonlinear attention
83 let nonlin_attention = NonlinAttention::load(dim, num_heads, vb.pp("nonlin_attention"), backend.clone())?;
84
85 // Two convolution modules
86 let conv_module1 = ConvolutionModule::load(dim, cnn_kernel, vb.pp("conv_module1"), backend.clone())?;
87 let conv_module2 = ConvolutionModule::load(dim, cnn_kernel, vb.pp("conv_module2"), backend)?;
88
89 // Bypass modules (per-channel scale)
90 let bypass = BypassModule::load_dim(dim, vb.pp("bypass"))?;
91 let bypass_mid = BypassModule::load_dim(dim, vb.pp("bypass_mid"))?;
92
93 Ok(Self {
94 norm,
95 feed_forward1,
96 feed_forward2,
97 feed_forward3,
98 self_attn_weights,
99 self_attn1,
100 self_attn2,
101 nonlin_attention,
102 conv_module1,

Callers

nothing calls this directly

Calls 1

cloneMethod · 0.45

Tested by

no test coverage detected