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

Method load

cake-core/src/models/flux/text_encoder.rs:43–86  ·  view source on GitHub ↗
(vb: VarBuilder, cfg: &EncoderConfig, backend: Arc<dyn ComputeBackend>)

Source from the content-addressed store, hash-verified

41
42impl EncoderBlock {
43 fn load(vb: VarBuilder, cfg: &EncoderConfig, backend: Arc<dyn ComputeBackend>) -> Result<Self> {
44 let h = cfg.hidden_size;
45 let i = cfg.intermediate_size;
46 let size_q = cfg.num_heads * cfg.head_dim;
47 let size_kv = cfg.num_kv_heads * cfg.head_dim;
48
49 let attn = vb.pp("self_attn");
50 let q_proj_weight = attn.pp("q_proj").get((size_q, h), "weight")?;
51 let k_proj_weight = attn.pp("k_proj").get((size_kv, h), "weight")?;
52 let v_proj_weight = attn.pp("v_proj").get((size_kv, h), "weight")?;
53 let o_proj_weight = attn.pp("o_proj").get((h, size_q), "weight")?;
54 let q_norm_weight = attn.pp("q_norm").get(cfg.head_dim, "weight")?;
55 let k_norm_weight = attn.pp("k_norm").get(cfg.head_dim, "weight")?;
56 let qk_norm_eps = cfg.rms_norm_eps as f32;
57
58 let mlp = vb.pp("mlp");
59 let gate_proj_weight = mlp.pp("gate_proj").get((i, h), "weight")?;
60 let up_proj_weight = mlp.pp("up_proj").get((i, h), "weight")?;
61 let down_proj_weight = mlp.pp("down_proj").get((h, i), "weight")?;
62
63 let rms_1_weight = vb.pp("input_layernorm").get(h, "weight")?;
64 let rms_2_weight = vb.pp("post_attention_layernorm").get(h, "weight")?;
65 let rms_eps = cfg.rms_norm_eps as f32;
66
67 Ok(Self {
68 rms_1_weight,
69 rms_2_weight,
70 rms_eps,
71 q_proj_weight,
72 k_proj_weight,
73 v_proj_weight,
74 o_proj_weight,
75 q_norm_weight,
76 k_norm_weight,
77 qk_norm_eps,
78 num_heads: cfg.num_heads,
79 num_kv_heads: cfg.num_kv_heads,
80 head_dim: cfg.head_dim,
81 gate_proj_weight,
82 up_proj_weight,
83 down_proj_weight,
84 backend,
85 })
86 }
87
88 /// attn_mask: optional (1, seq) tensor with 1 for real tokens, 0 for padding
89 fn forward_with_mask(&self, x: &Tensor, attn_mask: Option<&Tensor>) -> Result<Tensor> {

Callers

nothing calls this directly

Calls 2

getMethod · 0.45
cloneMethod · 0.45

Tested by

no test coverage detected