| 155 | {input.shape.dims[0], input.shape.dims[1], config.hidden_size})); |
| 156 | return o_proj.build(ctx, context, weights.o_proj); |
| 157 | } |
| 158 | |
| 159 | core::TensorValue feed_forward( |
| 160 | core::ModuleBuildContext & ctx, |
| 161 | const core::TensorValue & input, |
| 162 | const T5BaseEncoderLayerWeights & weights, |
| 163 | const T5BaseEncoderConfig & config) { |
| 164 | auto hidden = LinearModule({config.hidden_size, config.intermediate_size, false, GGML_PREC_F32}) |
| 165 | .build(ctx, input, weights.wi_proj); |
| 166 | if (config.feed_forward_kind == T5BaseFeedForwardKind::Relu) { |
| 167 | hidden = ReluModule{}.build(ctx, hidden); |
| 168 | } else if (config.feed_forward_kind == T5BaseFeedForwardKind::GatedGeluTanh) { |
| 169 | if (!weights.gate_proj.weight.valid()) { |
| 170 | throw std::runtime_error("T5BaseEncoder gated GELU feed-forward requires gate_proj.weight"); |
| 171 | } |
| 172 | auto gate = LinearModule({config.hidden_size, config.intermediate_size, false, GGML_PREC_F32}) |
| 173 | .build(ctx, input, weights.gate_proj); |
no test coverage detected