Skip to content

Commit b153194

Browse files
committed
Add ACE-Step caption rewrite action
1 parent 17751c0 commit b153194

6 files changed

Lines changed: 180 additions & 14 deletions

File tree

include/engine/models/ace_step/types.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,7 @@ struct AceStepRequest {
8585
std::optional<float> repainting_end_seconds = std::nullopt;
8686
std::optional<runtime::AudioBuffer> source_audio = std::nullopt;
8787
std::optional<AceStepReferenceCondition> reference = std::nullopt;
88+
bool rewrite_caption = false;
8889
AceStepGenerationOptions generation;
8990
};
9091

src/models/ace_step/request_parser.cpp

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,10 @@ AceStepRequest ace_step_parse_request(const runtime::TaskRequest &request) {
126126
if (const auto lyrics = runtime::find_option(request.options, {"lyrics"}); lyrics.has_value()) {
127127
out.lyrics = *lyrics;
128128
}
129+
if (const auto rewrite_caption = runtime::find_option(request.options, {"rewrite_caption"});
130+
rewrite_caption.has_value()) {
131+
out.rewrite_caption = runtime::parse_bool_option(*rewrite_caption, "rewrite_caption");
132+
}
129133
if (const auto negative = runtime::find_option(request.options, {"negative_prompt"}); negative.has_value()) {
130134
out.negative_prompt = *negative;
131135
}

src/models/ace_step/session.cpp

Lines changed: 78 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,16 +2,19 @@
22

33
#include "engine/framework/assets/tensor_source.h"
44
#include "engine/framework/debug/profiler.h"
5+
#include "engine/framework/io/json.h"
56
#include "engine/framework/runtime/options.h"
67
#include "engine/models/ace_step/prompt_builder.h"
78
#include "engine/models/ace_step/repaint.h"
89
#include "engine/models/ace_step/request_parser.h"
910
#include "engine/models/ace_step/task_route.h"
1011

1112
#include <chrono>
13+
#include <sstream>
1214
#include <stdexcept>
1315
#include <string>
1416
#include <string_view>
17+
#include <unordered_map>
1518
#include <utility>
1619

1720
namespace engine::models::ace_step {
@@ -143,6 +146,39 @@ bool mem_saver_from_options(const runtime::SessionOptions & options) {
143146
return false;
144147
}
145148

149+
bool rewrite_caption_requested(const std::unordered_map<std::string, std::string> & options) {
150+
const auto value = runtime::find_option(options, {"rewrite_caption"});
151+
return value.has_value() && runtime::parse_bool_option(*value, "rewrite_caption");
152+
}
153+
154+
std::string plan_json(const AceStepPlan & plan, const AceStepRequest & request) {
155+
std::ostringstream out;
156+
out << "{"
157+
<< "\"caption\":" << engine::io::json::stringify_string(plan.caption)
158+
<< ",\"cot_caption\":" << engine::io::json::stringify_string(plan.cot_caption)
159+
<< ",\"lyrics\":" << engine::io::json::stringify_string(request.lyrics);
160+
if (plan.metadata.bpm.has_value()) {
161+
out << ",\"bpm\":" << *plan.metadata.bpm;
162+
}
163+
if (plan.metadata.duration.has_value()) {
164+
out << ",\"duration_seconds\":" << *plan.metadata.duration;
165+
}
166+
if (plan.metadata.keyscale.has_value()) {
167+
out << ",\"keyscale\":" << engine::io::json::stringify_string(*plan.metadata.keyscale);
168+
}
169+
if (plan.metadata.timesignature.has_value()) {
170+
out << ",\"timesignature\":" << engine::io::json::stringify_string(*plan.metadata.timesignature);
171+
}
172+
if (plan.metadata.language.has_value()) {
173+
out << ",\"language\":" << engine::io::json::stringify_string(*plan.metadata.language);
174+
}
175+
if (plan.metadata.genres.has_value()) {
176+
out << ",\"genres\":" << engine::io::json::stringify_string(*plan.metadata.genres);
177+
}
178+
out << "}";
179+
return out.str();
180+
}
181+
146182
} // namespace
147183

148184
AceStepSession::AceStepSession(runtime::TaskSpec task, runtime::SessionOptions options,
@@ -178,8 +214,11 @@ runtime::RunMode AceStepSession::run_mode() const {
178214
}
179215

180216
void AceStepSession::prepare(const runtime::SessionPreparationRequest &request) {
181-
(void)request;
182217
ensure_planner();
218+
if (rewrite_caption_requested(request.options)) {
219+
mark_prepared();
220+
return;
221+
}
183222
ensure_vae_decoder();
184223
ensure_pre_dit();
185224
pre_dit_->prepare_runtime();
@@ -196,9 +235,46 @@ runtime::TaskResult AceStepSession::run(const runtime::TaskRequest &request) {
196235
const AceStepTaskRoute &route = ace_step_task_route(ace_request);
197236
validate_task_route_request(ace_request, route);
198237
engine::debug::timing_log_scalar("ace_step.session.parse_request_ms", engine::debug::elapsed_ms(parse_start, Clock::now()));
238+
const bool flow_edit_morph = ace_step_request_uses_flow_edit_morph(ace_request);
239+
if (ace_request.rewrite_caption) {
240+
if (flow_edit_morph || !route.uses_planner) {
241+
throw std::runtime_error("ACE-Step rewrite_caption requires a planner route");
242+
}
243+
const auto planner_ensure_start = Clock::now();
244+
ensure_planner();
245+
engine::debug::timing_log_scalar("ace_step.session.ensure_planner_ms",
246+
engine::debug::elapsed_ms(planner_ensure_start, Clock::now()));
247+
const auto planner_start = Clock::now();
248+
AceStepPlan plan = planner_->generate(ace_request, false);
249+
engine::debug::timing_log_scalar("ace_step.session.planner_generate_ms",
250+
engine::debug::elapsed_ms(planner_start, Clock::now()));
251+
const auto planner_release_start = Clock::now();
252+
planner_->release_graph_workspace();
253+
if (execution_context().backend_type() == core::BackendType::Metal) {
254+
planner_.reset();
255+
}
256+
engine::debug::timing_log_scalar("ace_step.session.planner_release.graph.workspace_ms",
257+
engine::debug::elapsed_ms(planner_release_start, Clock::now()));
258+
259+
runtime::TaskResult result;
260+
result.text_output = runtime::Transcript{
261+
plan.caption,
262+
plan.metadata.language.value_or(ace_request.vocal_language),
263+
};
264+
result.output_artifacts.push_back(runtime::make_text_artifact(
265+
runtime::ArtifactKind::Custom,
266+
"ace_step_caption_plan",
267+
plan_json(plan, ace_request),
268+
{
269+
{"format", "json"},
270+
{"extension", "json"},
271+
{"mime", "application/json"},
272+
}));
273+
engine::debug::timing_log_scalar("session.wall_ms", engine::debug::elapsed_ms(total_start, Clock::now()));
274+
return result;
275+
}
199276

200277
AceStepPlan plan;
201-
const bool flow_edit_morph = ace_step_request_uses_flow_edit_morph(ace_request);
202278
const bool has_request_audio_codes = !flow_edit_morph && !ace_request.audio_code_ids.empty();
203279
const bool use_planner =
204280
!flow_edit_morph &&

webui/native/dist/index.html

Lines changed: 10 additions & 10 deletions
Large diffs are not rendered by default.

webui/native/src/lib/i18n.ts

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -161,6 +161,9 @@ const english: Record<string, string> = {
161161
'request.randomSeed': '-1 = random',
162162
'request.maxTokens': 'Maximum tokens',
163163
'request.duration': 'Duration seconds',
164+
'request.autoDuration': '-1 = auto',
165+
'request.rewriteCaption': 'Rewrite caption',
166+
'request.rewritingCaption': 'Rewriting caption...',
164167
'request.minimaxFrames': '{frames} aligned output frames',
165168
'request.sourceAudio': 'Source audio',
166169
'request.stopRecording': 'Stop recording',

webui/native/src/routes/+page.svelte

Lines changed: 84 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,7 @@
6161
let installed: boolean | null = null;
6262
let loadingModel = false;
6363
let running = false;
64+
let rewritingCaption = false;
6465
let status = 'Ready';
6566
let warningStatus = '';
6667
let errorStatus = '';
@@ -230,7 +231,8 @@
230231
}
231232
232233
function setDuration(value: number) {
233-
duration = Math.max(1, Number.isFinite(value) ? value : 1);
234+
const minimum = selected?.family === 'ace_step' ? -1 : 1;
235+
duration = Math.max(minimum, Number.isFinite(value) ? value : minimum);
234236
if (selected?.family === 'minimax_h3') {
235237
advancedValues = { ...advancedValues, num_frames: miniMaxFramesForDuration(duration) };
236238
}
@@ -386,6 +388,7 @@
386388
modelMatchesSelectedPackage(model, selected));
387389
$: isFireRedAudioEdit = selected?.id === 'firered-audio-semantic-edit' ||
388390
selected?.id === 'firered-audio-acoustic-edit';
391+
$: allowsAutoDuration = selected?.family === 'ace_step';
389392
$: usesDurationSecOption =
390393
selected?.family === 'controlfoley' ||
391394
selected?.family === 'midashenglm_gen';
@@ -1220,6 +1223,74 @@
12201223
return { ...defaults, ...advancedValues, ...raw };
12211224
}
12221225
1226+
function base64Text(value: string): string {
1227+
const binary = atob(value);
1228+
const bytes = new Uint8Array(binary.length);
1229+
for (let index = 0; index < binary.length; index += 1) {
1230+
bytes[index] = binary.charCodeAt(index);
1231+
}
1232+
return new TextDecoder().decode(bytes);
1233+
}
1234+
1235+
function acePlanFromResult(result: Record<string, unknown>): Record<string, unknown> {
1236+
if (Array.isArray(result.artifacts)) {
1237+
const artifact = result.artifacts.find((entry): entry is { id: string; payload: string } =>
1238+
typeof entry === 'object' && entry !== null &&
1239+
(entry as { id?: unknown }).id === 'ace_step_caption_plan' &&
1240+
typeof (entry as { payload?: unknown }).payload === 'string');
1241+
if (artifact) return JSON.parse(base64Text(artifact.payload));
1242+
}
1243+
if (typeof result.text === 'string') return { caption: result.text };
1244+
return {};
1245+
}
1246+
1247+
async function rewriteAceCaption() {
1248+
if (selected?.family !== 'ace_step' || running || rewritingCaption) return;
1249+
if (!text.trim() && !lyrics.trim()) {
1250+
status = 'Enter a caption or lyrics to rewrite.';
1251+
warningStatus = status;
1252+
errorStatus = '';
1253+
return;
1254+
}
1255+
rewritingCaption = true;
1256+
warningStatus = '';
1257+
errorStatus = '';
1258+
status = tr('request.rewritingCaption');
1259+
try {
1260+
await ensureLoaded();
1261+
const options = { ...requestOptions(), rewrite_caption: true };
1262+
const request: Record<string, unknown> = {
1263+
text,
1264+
seed: resolveRequestSeed(seed),
1265+
duration_seconds: duration,
1266+
options
1267+
};
1268+
if (language.trim()) request.language = language;
1269+
if (lyrics.trim()) request.lyrics = lyrics;
1270+
const result = await runTask({ model: selected.id, request });
1271+
const plan = acePlanFromResult(result);
1272+
if (typeof plan.caption === 'string' && plan.caption.trim()) text = plan.caption;
1273+
if (typeof plan.language === 'string' && plan.language.trim()) language = plan.language;
1274+
if (typeof plan.duration_seconds === 'number' && Number.isFinite(plan.duration_seconds) && plan.duration_seconds > 0) {
1275+
duration = plan.duration_seconds;
1276+
}
1277+
const nextAdvanced = { ...advancedValues };
1278+
if (typeof plan.bpm === 'number' && Number.isFinite(plan.bpm)) nextAdvanced.bpm = plan.bpm;
1279+
if (typeof plan.keyscale === 'string') nextAdvanced.keyscale = plan.keyscale;
1280+
if (typeof plan.timesignature === 'string') nextAdvanced.timesignature = plan.timesignature;
1281+
advancedValues = nextAdvanced;
1282+
outputText = typeof result.text === 'string' ? result.text : '';
1283+
outputJson = JSON.stringify(plan, null, 2);
1284+
status = 'Caption rewritten.';
1285+
} catch (error) {
1286+
status = error instanceof Error ? error.message : String(error);
1287+
errorStatus = status;
1288+
log(`Caption rewrite failed: ${status}`);
1289+
} finally {
1290+
rewritingCaption = false;
1291+
}
1292+
}
1293+
12231294
function clearOutput() {
12241295
for (const output of outputAudio) URL.revokeObjectURL(output.url);
12251296
outputAudio = [];
@@ -2112,6 +2183,14 @@
21122183
<label for="lyrics">{tr('request.lyrics')} <span>{lyricsRequired ? tr('voice.required') : tr('request.optional')}</span></label>
21132184
<textarea id="lyrics" rows="3" bind:value={lyrics} required={lyricsRequired}
21142185
aria-required={lyricsRequired} placeholder="[Verse]…"></textarea>
2186+
{#if selected.family === 'ace_step'}
2187+
<div class="media-actions">
2188+
<button type="button" disabled={running || rewritingCaption || (!text.trim() && !lyrics.trim())}
2189+
on:click={rewriteAceCaption}>
2190+
{rewritingCaption ? tr('request.rewritingCaption') : tr('request.rewriteCaption')}
2191+
</button>
2192+
</div>
2193+
{/if}
21152194
{/if}
21162195

21172196
{#if selected.task === 'asr'}
@@ -2147,8 +2226,11 @@
21472226
{#if selected.task === 'gen'}
21482227
<div>
21492228
<label for="duration">{tr('request.duration')}</label>
2150-
<input id="duration" type="number" min="1" step="0.1" value={duration}
2229+
<input id="duration" type="number" min={allowsAutoDuration ? -1 : 1} step="0.1" value={duration}
21512230
on:input={(event) => setDuration(event.currentTarget.valueAsNumber)} />
2231+
{#if allowsAutoDuration}
2232+
<small>{tr('request.autoDuration')}</small>
2233+
{/if}
21522234
{#if selected.family === 'minimax_h3'}
21532235
<small>{tr('request.minimaxFrames', { frames: Number(advancedValues.num_frames || 0) })}</small>
21542236
{/if}

0 commit comments

Comments
 (0)