Repository navigation
Expand file tree
/
Copy pathbatch.cpp
More file actions
300 lines (277 loc) · 10.8 KB
/
Copy pathbatch.cpp
File metadata and controls
300 lines (277 loc) · 10.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
#include "batch.h"
#include "args.h"
#include "request.h"
#include "engine/framework/io/json.h"
#include <algorithm>
#include <cctype>
#include <fstream>
#include <iterator>
#include <sstream>
#include <stdexcept>
#include <utility>
namespace minitts::cli {
namespace {
bool is_wav_path(const std::filesystem::path & path) {
std::string ext = path.extension().string();
std::transform(ext.begin(), ext.end(), ext.begin(), [](unsigned char ch) {
return static_cast<char>(std::tolower(ch));
});
return ext == ".wav";
}
std::string lowercase_extension(const std::filesystem::path & path) {
std::string ext = path.extension().string();
std::transform(ext.begin(), ext.end(), ext.begin(), [](unsigned char ch) {
return static_cast<char>(std::tolower(ch));
});
return ext;
}
std::string read_text_file(const std::filesystem::path & path) {
std::ifstream input(path);
if (!input) {
throw std::runtime_error("failed to open batch text file: " + path.string());
}
std::ostringstream raw;
raw << input.rdbuf();
return raw.str();
}
std::string normalize_text_as_paragraph(const std::string & text) {
std::string normalized;
bool in_space = false;
for (unsigned char ch : text) {
if (std::isspace(ch)) {
in_space = !normalized.empty();
continue;
}
if (in_space) {
normalized.push_back(' ');
in_space = false;
}
normalized.push_back(static_cast<char>(ch));
}
return normalized;
}
std::string read_plain_text_batch_file(const std::filesystem::path & path) {
return normalize_text_as_paragraph(read_text_file(path));
}
std::string read_json_text_batch_file(const std::filesystem::path & path) {
const auto root = engine::io::json::parse_file(path);
if (root.is_string()) {
return normalize_text_as_paragraph(root.as_string());
}
if (root.is_object()) {
if (const auto * input = root.find("input"); input != nullptr && input->is_string()) {
return normalize_text_as_paragraph(input->as_string());
}
if (const auto * text = root.find("text"); text != nullptr && text->is_string()) {
return normalize_text_as_paragraph(text->as_string());
}
}
throw std::runtime_error("--batch-text-dir JSON file requires a string root, input, or text: " + path.string());
}
struct TextBatchFileFormat {
const char * extension;
std::string (*read)(const std::filesystem::path & path);
};
const TextBatchFileFormat kTextBatchFileFormats[] = {
{".txt", read_plain_text_batch_file},
{".md", read_plain_text_batch_file},
{".json", read_json_text_batch_file},
};
const TextBatchFileFormat * text_batch_format_for(const std::filesystem::path & path) {
const auto ext = lowercase_extension(path);
for (const auto & format : kTextBatchFileFormats) {
if (ext == format.extension) {
return &format;
}
}
return nullptr;
}
std::string supported_text_batch_extensions() {
std::ostringstream out;
for (size_t i = 0; i < std::size(kTextBatchFileFormats); ++i) {
if (i != 0) {
out << ", ";
}
out << kTextBatchFileFormats[i].extension;
}
return out.str();
}
minitts::app::AppBatchRequest build_request_sequence_from_json(
const std::filesystem::path & sequence_path) {
const auto root = engine::io::json::parse_file(sequence_path);
const auto * requests_value = root.is_array() ? &root : root.find("requests");
if (requests_value == nullptr) {
throw std::runtime_error("request sequence json requires an array or a requests array");
}
minitts::app::AppBatchRequest batch;
const auto base_dir = sequence_path.parent_path();
int index = 0;
for (const auto & item : requests_value->as_array()) {
std::ostringstream fallback;
fallback << "request_" << index;
batch.requests.push_back(minitts::app::AppRequest{
json_optional_string(item, "id").value_or(fallback.str()),
build_request_from_json(item, base_dir),
});
++index;
}
if (batch.requests.empty()) {
throw std::runtime_error("request sequence must contain at least one request");
}
return batch;
}
minitts::app::AppBatchRequest build_text_file_batch(
const std::filesystem::path & path,
const engine::runtime::TaskRequest & base_request,
const std::string & language) {
std::ifstream input(path);
if (!input) {
throw std::runtime_error("failed to open --batch-text-file: " + path.string());
}
minitts::app::AppBatchRequest batch;
std::string line;
int line_number = 0;
while (std::getline(input, line)) {
++line_number;
if (!line.empty() && line.back() == '\r') {
line.pop_back();
}
if (line.empty()) {
continue;
}
auto request = base_request;
request.text_input = engine::runtime::Transcript{line, language};
std::ostringstream id;
id << "line_" << line_number;
batch.requests.push_back(minitts::app::AppRequest{id.str(), std::move(request)});
}
if (batch.requests.empty()) {
throw std::runtime_error("--batch-text-file contains no non-empty requests");
}
return batch;
}
minitts::app::AppBatchRequest build_text_dir_batch(
const std::filesystem::path & dir,
const engine::runtime::TaskRequest & base_request,
const std::string & language) {
if (!std::filesystem::is_directory(dir)) {
throw std::runtime_error("--batch-text-dir must be an existing directory: " + dir.string());
}
std::vector<std::filesystem::path> text_paths;
for (const auto & entry : std::filesystem::directory_iterator(dir)) {
if (entry.is_regular_file() && text_batch_format_for(entry.path()) != nullptr) {
text_paths.push_back(entry.path());
}
}
std::sort(text_paths.begin(), text_paths.end());
if (text_paths.empty()) {
throw std::runtime_error(
"--batch-text-dir contains no supported text files (" + supported_text_batch_extensions() + "): " +
dir.string());
}
minitts::app::AppBatchRequest batch;
batch.requests.reserve(text_paths.size());
for (const auto & path : text_paths) {
const auto * format = text_batch_format_for(path);
const auto paragraph = format->read(path);
if (paragraph.empty()) {
throw std::runtime_error("--batch-text-dir file contains no non-whitespace text: " + path.string());
}
auto request = base_request;
request.text_input = engine::runtime::Transcript{paragraph, language};
batch.requests.push_back(minitts::app::AppRequest{path.stem().string(), std::move(request)});
}
return batch;
}
minitts::app::AppBatchRequest build_audio_dir_batch(
const std::filesystem::path & dir,
const engine::runtime::TaskRequest & base_request,
const std::string & audio_role) {
if (!std::filesystem::is_directory(dir)) {
throw std::runtime_error("--batch-audio-dir must be an existing directory: " + dir.string());
}
std::vector<std::filesystem::path> audio_paths;
for (const auto & entry : std::filesystem::directory_iterator(dir)) {
if (entry.is_regular_file() && is_wav_path(entry.path())) {
audio_paths.push_back(entry.path());
}
}
std::sort(audio_paths.begin(), audio_paths.end());
minitts::app::AppBatchRequest batch;
batch.requests.reserve(audio_paths.size());
for (const auto & path : audio_paths) {
auto request = base_request;
if (audio_role == "audio") {
request.audio_input = read_audio_buffer(path);
} else if (audio_role == "voice_ref") {
if (!request.voice.has_value()) {
request.voice = engine::runtime::VoiceCondition{};
}
if (!request.voice->speaker.has_value()) {
request.voice->speaker = engine::runtime::VoiceReference{};
}
request.voice->speaker->audio = read_audio_buffer(path);
} else if (audio_role == "source_audio" ||
audio_role == "target_voice" ||
audio_role == "prosody_ref" ||
audio_role == "style_ref") {
set_option(request.options, audio_role, path.string());
} else {
throw std::runtime_error(
"--batch-audio-role must be audio, voice_ref, source_audio, target_voice, prosody_ref, or style_ref");
}
batch.requests.push_back(minitts::app::AppRequest{path.stem().string(), std::move(request)});
}
if (batch.requests.empty()) {
throw std::runtime_error("--batch-audio-dir contains no .wav files: " + dir.string());
}
return batch;
}
} // namespace
bool has_batch_input(int argc, char ** argv) {
int count = 0;
count += optional_path_arg(argc, argv, "--request-sequence").has_value() ? 1 : 0;
count += optional_path_arg(argc, argv, "--batch-text-file").has_value() ? 1 : 0;
count += optional_path_arg(argc, argv, "--batch-text-dir").has_value() ? 1 : 0;
count += optional_path_arg(argc, argv, "--batch-audio-dir").has_value() ? 1 : 0;
return count > 0;
}
minitts::app::AppBatchRequest build_batch_request_from_cli(
int argc,
char ** argv,
const engine::runtime::TaskRequest & base_request,
const std::string & audio_role) {
const auto request_sequence_path = optional_path_arg(argc, argv, "--request-sequence");
const auto batch_text_file = optional_path_arg(argc, argv, "--batch-text-file");
const auto batch_text_dir = optional_path_arg(argc, argv, "--batch-text-dir");
const auto batch_audio_dir = optional_path_arg(argc, argv, "--batch-audio-dir");
const int count =
(request_sequence_path.has_value() ? 1 : 0) +
(batch_text_file.has_value() ? 1 : 0) +
(batch_text_dir.has_value() ? 1 : 0) +
(batch_audio_dir.has_value() ? 1 : 0);
if (count == 0) {
return minitts::app::AppBatchRequest{};
}
if (count > 1) {
throw std::runtime_error(
"choose only one of --request-sequence, --batch-text-file, --batch-text-dir, or --batch-audio-dir");
}
if (request_sequence_path.has_value()) {
return build_request_sequence_from_json(*request_sequence_path);
}
if (batch_text_file.has_value()) {
return build_text_file_batch(
*batch_text_file,
base_request,
find_arg(argc, argv, "--language").value_or(""));
}
if (batch_text_dir.has_value()) {
return build_text_dir_batch(
*batch_text_dir,
base_request,
find_arg(argc, argv, "--language").value_or(""));
}
return build_audio_dir_batch(*batch_audio_dir, base_request, audio_role);
}
} // namespace minitts::cli