MCPcopy Create free account
hub / github.com/cactus-compute/cactus / compute_stft_node

Function compute_stft_node

cactus/graph/graph_ops_nn.cpp:1817–1837  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1815}
1816
1817void 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
1839void 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) {

Callers

nothing calls this directly

Calls 1

cactus_stft_f16Function · 0.85

Tested by

no test coverage detected