| 1815 | } |
| 1816 | |
| 1817 | void compute_stft_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, |
| 1818 | const std::unordered_map<size_t, size_t>& node_index_map) { |
| 1819 | const auto& X = get_input(node, 0, nodes, node_index_map); |
| 1820 | const auto& W = get_input(node, 1, nodes, node_index_map); |
| 1821 | auto& Y = node.output_buffer; |
| 1822 | |
| 1823 | const size_t N = X.shape[0]; |
| 1824 | const size_t C_in = X.shape[1]; |
| 1825 | const size_t L = X.shape[2]; |
| 1826 | const size_t C_out = W.shape[0]; |
| 1827 | const size_t K = W.shape[2]; |
| 1828 | const size_t stride = node.params.stride; |
| 1829 | const size_t num_fft_bins = node.params.num_fft_bins; |
| 1830 | |
| 1831 | if (X.precision != Precision::FP16 || W.precision != Precision::FP16) { |
| 1832 | throw std::runtime_error("stft only supports FP16"); |
| 1833 | } |
| 1834 | |
| 1835 | cactus_stft_f16(X.data_as<__fp16>(), W.data_as<__fp16>(), |
| 1836 | Y.data_as<__fp16>(), N, L, C_in, C_out, K, stride, num_fft_bins); |
| 1837 | } |
| 1838 | |
| 1839 | void compute_conv1d_k7s3_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, |
| 1840 | const std::unordered_map<size_t, size_t>& node_index_map) { |
nothing calls this directly
no test coverage detected