| 123 | throw std::runtime_error("ModuleBuildContext.ggml is null"); |
| 124 | } |
| 125 | core::validate_rank_between(input, 3, 3, "input"); |
| 126 | core::validate_shape( |
| 127 | input, |
| 128 | core::TensorShape::from_dims({input.shape.dims[0], config_.channels, input.shape.dims[2]}), |
| 129 | "input"); |
| 130 | core::validate_shape( |
| 131 | weights.weight, |
| 132 | core::TensorShape::from_dims({config_.channels, 1, config_.kernel_size}), |
| 133 | "weight"); |
| 134 | const auto input_contiguous = ensure_f32(ctx, tensor_layout::ensure_contiguous_layout_if_needed(ctx, input)); |
| 135 | const auto weight_contiguous = regular_conv_weight(ctx, weights.weight, "DepthwiseConv1dModule"); |
| 136 | auto input_4d = core::reshape_tensor( |
| 137 | ctx, |
| 138 | input_contiguous, |
| 139 | core::TensorShape::from_dims({input.shape.dims[0], config_.channels, 1, input.shape.dims[2]})); |
| 140 | auto weight_4d = core::reshape_tensor( |
| 141 | ctx, |
| 142 | weight_contiguous, |
| 143 | core::TensorShape::from_dims({config_.channels, 1, 1, config_.kernel_size})); |
| 144 | // This 2D depthwise lowering greatly improves performance but may affect parity. |
| 145 | auto output_4d = DepthwiseConv2dModule({ |
| 146 | config_.channels, |
| 147 | 1, |
| 148 | config_.kernel_size, |
| 149 | 1, |
| 150 | config_.stride, |
| 151 | 0, |
| 152 | config_.padding, |
| 153 | 1, |
| 154 | config_.dilation, |
| 155 | config_.use_bias, |
| 156 | }).build(ctx, input_4d, {weight_4d, weights.bias}); |
| 157 | return core::reshape_tensor( |
| 158 | ctx, |
| 159 | output_4d, |
| 160 | core::TensorShape::from_dims({ |
| 161 | input.shape.dims[0], |
| 162 | config_.channels, |
| 163 | depthwise_conv1d_output_frames(config_, input.shape.dims[2]), |
| 164 | })); |
| 165 | } |
| 166 |
no test coverage detected