diff --git a/otherarch/sdcpp/examples/common/media_io.cpp b/otherarch/sdcpp/examples/common/media_io.cpp index 506c67f4d..f0fdf374e 100644 --- a/otherarch/sdcpp/examples/common/media_io.cpp +++ b/otherarch/sdcpp/examples/common/media_io.cpp @@ -191,6 +191,20 @@ uint32_t read_u32_le_bytes(const uint8_t* data) { (static_cast(data[3]) << 24); } +uint16_t read_u16_le_bytes(const uint8_t* p) { + return static_cast(p[0]) | (static_cast(p[1]) << 8); +} + +int32_t read_s24_le_bytes(const uint8_t* p) { + int32_t value = static_cast(p[0]) | + (static_cast(p[1]) << 8) | + (static_cast(p[2]) << 16); + if (value & 0x00800000) { + value |= 0xff000000; + } + return value; +} + int stbi_ext_write_png_to_func(stbi_write_func* func, void* context, int x, @@ -1374,3 +1388,104 @@ bool write_wav_to_file(const std::string& path, file.write(reinterpret_cast(pcm.data()), static_cast(pcm.size() * sizeof(int16_t))); return file.good(); } + +sd_audio_t load_pcm_wav_from_file(const std::string& path) { + sd_audio_t audio = {0, 0, 0, nullptr}; + if (path.empty()) { + return audio; + } + + std::vector wav; + if (!read_binary_file_bytes(path.c_str(), wav)) { + LOG_ERROR("load WAV from '%s' failed", path.c_str()); + return audio; + } + if (wav.size() < 44 || std::memcmp(wav.data(), "RIFF", 4) != 0 || std::memcmp(wav.data() + 8, "WAVE", 4) != 0) { + LOG_ERROR("input audio file '%s' is not a RIFF/WAVE file", path.c_str()); + return audio; + } + + uint16_t format = 0; + uint16_t channels = 0; + uint32_t sample_rate = 0; + uint16_t bits_per_sample = 0; + const uint8_t* data = nullptr; + uint32_t data_size = 0; + + size_t pos = 12; + while (pos + 8 <= wav.size()) { + const uint8_t* chunk = wav.data() + pos; + uint32_t chunk_size = read_u32_le_bytes(chunk + 4); + size_t chunk_data = pos + 8; + if (chunk_data + chunk_size > wav.size()) { + break; + } + + if (std::memcmp(chunk, "fmt ", 4) == 0 && chunk_size >= 16) { + format = read_u16_le_bytes(wav.data() + chunk_data); + channels = read_u16_le_bytes(wav.data() + chunk_data + 2); + sample_rate = read_u32_le_bytes(wav.data() + chunk_data + 4); + bits_per_sample = read_u16_le_bytes(wav.data() + chunk_data + 14); + } else if (std::memcmp(chunk, "data", 4) == 0) { + data = wav.data() + chunk_data; + data_size = chunk_size; + } + pos = chunk_data + chunk_size + (chunk_size & 1); + } + + if (data == nullptr || data_size == 0 || channels == 0 || sample_rate == 0) { + LOG_ERROR("input WAV '%s' is missing fmt/data chunks", path.c_str()); + return audio; + } + if (format != 1 && format != 3) { + LOG_ERROR("unsupported WAV format %u in '%s', only PCM and float WAV are supported", + static_cast(format), + path.c_str()); + return audio; + } + + uint16_t bytes_per_sample = static_cast((bits_per_sample + 7) / 8); + uint32_t frame_bytes = static_cast(bytes_per_sample) * channels; + if (bytes_per_sample == 0 || frame_bytes == 0 || data_size < frame_bytes) { + LOG_ERROR("invalid WAV sample format in '%s'", path.c_str()); + return audio; + } + + uint64_t sample_count = data_size / frame_bytes; + size_t float_count = static_cast(sample_count) * channels; + float* samples = (float*)malloc(float_count * sizeof(float)); + if (samples == nullptr) { + return audio; + } + + for (uint64_t i = 0; i < sample_count; ++i) { + for (uint16_t ch = 0; ch < channels; ++ch) { + const uint8_t* src = data + i * frame_bytes + ch * bytes_per_sample; + float sample = 0.f; + if (format == 3 && bits_per_sample == 32) { + std::memcpy(&sample, src, sizeof(float)); + } else if (format == 1 && bits_per_sample == 8) { + sample = (static_cast(src[0]) - 128) / 128.f; + } else if (format == 1 && bits_per_sample == 16) { + sample = static_cast(read_u16_le_bytes(src)) / 32768.f; + } else if (format == 1 && bits_per_sample == 24) { + sample = read_s24_le_bytes(src) / 8388608.f; + } else if (format == 1 && bits_per_sample == 32) { + sample = static_cast(read_u32_le_bytes(src)) / 2147483648.f; + } else { + LOG_ERROR("unsupported WAV bit depth %u in '%s'", + static_cast(bits_per_sample), + path.c_str()); + free(samples); + return audio; + } + samples[i * channels + ch] = std::clamp(sample, -1.0f, 1.0f); + } + } + + audio.sample_rate = sample_rate; + audio.channels = channels; + audio.sample_count = sample_count; + audio.data = samples; + return audio; +} diff --git a/otherarch/sdcpp/examples/common/media_io.h b/otherarch/sdcpp/examples/common/media_io.h index 0f7679d7f..df2fd019b 100644 --- a/otherarch/sdcpp/examples/common/media_io.h +++ b/otherarch/sdcpp/examples/common/media_io.h @@ -110,4 +110,6 @@ bool write_wav_to_file(const std::string& path, uint32_t channels, uint32_t sample_rate); +sd_audio_t load_pcm_wav_from_file(const std::string& path); + #endif // __MEDIA_IO_H__ diff --git a/otherarch/sdcpp/src/model/vae/ltx_audio_vae.hpp b/otherarch/sdcpp/src/model/vae/ltx_audio_vae.hpp index 7a3fb25a2..49847445d 100644 --- a/otherarch/sdcpp/src/model/vae/ltx_audio_vae.hpp +++ b/otherarch/sdcpp/src/model/vae/ltx_audio_vae.hpp @@ -250,9 +250,9 @@ namespace LTXV { sd::Tensor basis({n_fft, 1, n_freqs * 2}); for (int k = 0; k < n_freqs; ++k) { for (int n = 0; n < n_fft; ++n) { - double window = 0.5 - 0.5 * std::cos(2.0 * kPi * n / static_cast(n_fft)); - double phase = 2.0 * kPi * k * n / static_cast(n_fft); - basis.index(n, 0, k) = static_cast(std::cos(phase) * window); + double window = 0.5 - 0.5 * std::cos(2.0 * kPi * n / static_cast(n_fft)); + double phase = 2.0 * kPi * k * n / static_cast(n_fft); + basis.index(n, 0, k) = static_cast(std::cos(phase) * window); basis.index(n, 0, k + n_freqs) = static_cast(-std::sin(phase) * window); } } @@ -276,12 +276,12 @@ namespace LTXV { } for (int m = 0; m < n_mels; ++m) { - double lower = mel_f[m]; + double lower = mel_f[m]; double center = mel_f[m + 1]; - double upper = mel_f[m + 2]; - double enorm = 2.0 / std::max(upper - lower, 1e-12); + double upper = mel_f[m + 2]; + double enorm = 2.0 / std::max(upper - lower, 1e-12); for (int f = 0; f < n_freqs; ++f) { - double freq = fft_freqs[f]; + double freq = fft_freqs[f]; double value = 0.0; if (freq > lower && freq <= center) { value = (freq - lower) / std::max(center - lower, 1e-12); diff --git a/otherarch/sdcpp/src/stable-diffusion.cpp b/otherarch/sdcpp/src/stable-diffusion.cpp index 765d59fd3..e2a80a035 100644 --- a/otherarch/sdcpp/src/stable-diffusion.cpp +++ b/otherarch/sdcpp/src/stable-diffusion.cpp @@ -3345,8 +3345,8 @@ static sd::Tensor sd_audio_to_ltx_waveform_tensor(const sd_audio_t* audio int64_t out_samples = static_cast(out_samples_u64); sd::Tensor waveform({out_samples, target_channels, 1, 1}); - const double src_rate = static_cast(audio->sample_rate); - const double dst_rate = static_cast(target_sample_rate); + const double src_rate = static_cast(audio->sample_rate); + const double dst_rate = static_cast(target_sample_rate); const int src_channels = static_cast(audio->channels); auto src_value = [&](uint64_t sample, int channel) -> float { @@ -3365,8 +3365,8 @@ static sd::Tensor sd_audio_to_ltx_waveform_tensor(const sd_audio_t* audio uint64_t i1 = std::min(i0 + 1, audio->sample_count - 1); float frac = static_cast(src_pos - static_cast(i0)); for (int ch = 0; ch < target_channels; ++ch) { - float v0 = src_value(i0, ch); - float v1 = src_value(i1, ch); + float v0 = src_value(i0, ch); + float v1 = src_value(i1, ch); waveform.index(t, ch, 0, 0) = v0 + (v1 - v0) * frac; } } @@ -5037,9 +5037,9 @@ static std::optional prepare_video_generation_latents(sd } int64_t audio_encode_start = ggml_time_ms(); - auto waveform = sd_audio_to_ltx_waveform_tensor(sd_vid_gen_params->input_audio, - sd_ctx->sd->audio_vae_model->config.sample_rate, - sd_ctx->sd->audio_vae_model->config.audio_channels); + auto waveform = sd_audio_to_ltx_waveform_tensor(sd_vid_gen_params->input_audio, + sd_ctx->sd->audio_vae_model->config.sample_rate, + sd_ctx->sd->audio_vae_model->config.audio_channels); if (waveform.empty()) { LOG_ERROR("failed to convert source audio for LTX A2V encoding"); return std::nullopt;