| 1571 | ~HTDemucsGraph() override = default; |
| 1572 | |
| 1573 | void run( |
| 1574 | const std::vector<float> & freq_input, |
| 1575 | const std::vector<float> & time_input, |
| 1576 | std::vector<float> & freq_output, |
| 1577 | std::vector<float> & time_output) { |
| 1578 | const size_t freq_elements = static_cast<size_t>( |
| 1579 | config_.input_freq_channels * config_.stft_freq_bins * config_.stft_frames); |
| 1580 | const size_t time_elements = static_cast<size_t>( |
| 1581 | config_.audio_channels * config_.segment_samples); |
| 1582 | if (freq_input.size() != freq_elements || time_input.size() != time_elements) { |
| 1583 | throw std::runtime_error("HTDemucs graph input size mismatch"); |
| 1584 | } |
| 1585 | rebuild(); |
| 1586 | const auto upload_start = std::chrono::steady_clock::now(); |
| 1587 | ggml_backend_tensor_set(freq_input_, freq_input.data(), 0, freq_input.size() * sizeof(float)); |
| 1588 | ggml_backend_tensor_set(time_input_, time_input.data(), 0, time_input.size() * sizeof(float)); |
| 1589 | const auto upload_end = std::chrono::steady_clock::now(); |
| 1590 | const auto compute_start = upload_end; |
| 1591 | core::set_backend_threads(backend_, threads_); |
| 1592 | const ggml_status status = engine::core::compute_backend_graph(backend_, graph_); |
| 1593 | ggml_backend_synchronize(backend_); |
| 1594 | const auto compute_end = std::chrono::steady_clock::now(); |
| 1595 | if (status != GGML_STATUS_SUCCESS) { |
| 1596 | throw std::runtime_error("HTDemucs graph compute failed"); |
| 1597 | } |
| 1598 | const auto readback_start = compute_end; |
| 1599 | if (freq_output.size() != static_cast<size_t>( |
| 1600 | config_.output_freq_channels * config_.stft_freq_bins * config_.stft_frames) || |
| 1601 | time_output.size() != static_cast<size_t>( |
| 1602 | config_.output_time_channels * config_.segment_samples)) { |
| 1603 | throw std::runtime_error("HTDemucs graph output buffer size mismatch"); |
| 1604 | } |
| 1605 | ggml_backend_tensor_get(freq_output_, freq_output.data(), 0, freq_output.size() * sizeof(float)); |
| 1606 | ggml_backend_tensor_get(time_output_, time_output.data(), 0, time_output.size() * sizeof(float)); |
| 1607 | const auto readback_end = std::chrono::steady_clock::now(); |
| 1608 | last_timing_.upload_ms = debug::elapsed_ms(upload_start, upload_end); |
| 1609 | last_timing_.compute_ms = debug::elapsed_ms(compute_start, compute_end); |
| 1610 | last_timing_.readback_ms = debug::elapsed_ms(readback_start, readback_end); |
| 1611 | } |
| 1612 | |
| 1613 | const Timing & last_timing() const noexcept { return last_timing_; } |
| 1614 |
no test coverage detected