Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions include/engine/models/irodori_tts/session.h
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@ class IrodoriTTSSession final : public runtime::RuntimeSessionBase,
assets::TensorStorageType codec_weight_storage_type_ =
assets::TensorStorageType::Native;
bool mem_saver_ = true;
std::unique_ptr<engine::core::ExecutionContext> codec_execution_context_;
std::unique_ptr<IrodoriConditionEncoder> condition_encoder_;
std::unique_ptr<IrodoriRfSampler> rf_sampler_;
std::unique_ptr<IrodoriCodec> codec_;
Expand Down
11 changes: 11 additions & 0 deletions model_specs/irodori_tts.json
Original file line number Diff line number Diff line change
Expand Up @@ -194,6 +194,17 @@
"required": false,
"default": "native"
},
{
"name": "codec_backend",
"type": "enum",
"description": "DACVAE codec execution backend. Use cpu as an opt-in workaround for backend-specific codec decoder issues; default same.",
"values": [
"same",
"cpu"
],
"required": false,
"default": "same"
},
{
"name": "condition_graph_arena_mb",
"type": "int",
Expand Down
44 changes: 42 additions & 2 deletions src/models/irodori_tts/session.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,11 @@ namespace {
using Clock = std::chrono::steady_clock;
constexpr const char *kFamily = "irodori_tts";

enum class IrodoriCodecBackend {
Same,
Cpu,
};

std::shared_ptr<const IrodoriTTSAssets>
require_assets(std::shared_ptr<const IrodoriTTSAssets> assets) {
if (assets == nullptr) {
Expand All @@ -45,6 +50,20 @@ std::shared_ptr<const engine::model_spec::ModelContract> require_contract(
return contract;
}

IrodoriCodecBackend parse_codec_backend(const runtime::SessionOptions & options) {
if (const auto value =
runtime::find_option(options.options, {"irodori_tts.codec_backend"})) {
if (*value == "same") {
return IrodoriCodecBackend::Same;
}
if (*value == "cpu") {
return IrodoriCodecBackend::Cpu;
}
throw std::runtime_error("Invalid irodori_tts.codec_backend: " + *value);
}
return IrodoriCodecBackend::Same;
}

runtime::SessionOptions normalize_session_options(runtime::SessionOptions options) {
return runtime::apply_option_v1_compatibility(
std::move(options),
Expand All @@ -60,8 +79,15 @@ runtime::SessionOptions require_supported_session_options(
const std::shared_ptr<const engine::model_spec::ModelContract> &contract) {
options = normalize_session_options(std::move(options));
const auto checked_contract = require_contract(contract);
auto validation_options = options;
// Older standalone GGUF packages embed a v1 contract that predates this
// workaround option; keep them usable while still validating the value below.
if (checked_contract->session_option_keys.find("irodori_tts.codec_backend") ==
checked_contract->session_option_keys.end()) {
validation_options.options.erase("irodori_tts.codec_backend");
}
runtime::validate_spec_backed_session_options(
options, *checked_contract, kFamily, "Irodori-TTS");
validation_options, *checked_contract, kFamily, "Irodori-TTS");
return options;
}

Expand Down Expand Up @@ -376,14 +402,24 @@ IrodoriTTSSession::IrodoriTTSSession(
throw std::runtime_error(
"Irodori-TTS supports only TTS, voice-cloning, and voice-design offline tasks");
}
const auto codec_backend = parse_codec_backend(this->options());
engine::core::ExecutionContext * codec_execution = &execution_context();
if (codec_backend == IrodoriCodecBackend::Cpu) {
auto codec_backend_config = this->options().backend;
codec_backend_config.type = engine::core::BackendType::Cpu;
codec_backend_config.device = 0;
codec_execution_context_ =
std::make_unique<engine::core::ExecutionContext>(codec_backend_config);
codec_execution = codec_execution_context_.get();
}
condition_encoder_ = std::make_unique<IrodoriConditionEncoder>(
assets_, execution_context(), condition_graph_arena_bytes_,
condition_weight_context_bytes_, weight_storage_type_);
rf_sampler_ = std::make_unique<IrodoriRfSampler>(
assets_, execution_context(), rf_graph_arena_bytes_,
rf_weight_context_bytes_, weight_storage_type_, mem_saver_);
codec_ = std::make_unique<IrodoriCodec>(
assets_, execution_context(), codec_graph_arena_bytes_,
assets_, *codec_execution, codec_graph_arena_bytes_,
codec_weight_context_bytes_, codec_weight_storage_type_);
assets_->model_weights->release_storage();
assets_->codec_weights->release_storage();
Expand All @@ -397,6 +433,10 @@ IrodoriTTSSession::IrodoriTTSSession(
assets_->config.max_text_len);
debug::trace_log_scalar("irodori_tts.config.max_caption_len",
assets_->config.max_caption_len);
debug::trace_log_scalar(
"irodori_tts.codec.backend",
std::string_view(codec_backend == IrodoriCodecBackend::Cpu ? "cpu"
: "same"));
}

IrodoriTTSSession::~IrodoriTTSSession() = default;
Expand Down
Loading