mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-10-03 03:25:40 +00:00
metadata forcing
This commit is contained in:
parent
444b9e4419
commit
60acd9b92f
1 changed files with 47 additions and 14 deletions
|
|
@ -589,10 +589,21 @@ struct MetadataFSM {
|
|||
value_acc.clear();
|
||||
}
|
||||
|
||||
// Force FSM to only allow a specific language value
|
||||
void force_language(BPETokenizer & bpe, const std::string & lang) {
|
||||
language_tree = PrefixTree();
|
||||
build_value_tree(bpe, language_tree, "language:", {lang});
|
||||
// Force a user-supplied metadata value. The constrained token sequence is
|
||||
// emitted into the CoT, so the LM sees it before generating lyrics/codes.
|
||||
void force_field(BPETokenizer & bpe, State value_state, const std::string & value) {
|
||||
PrefixTree * tree = nullptr;
|
||||
const char * prefix = nullptr;
|
||||
switch (value_state) {
|
||||
case BPM_VALUE: tree = &bpm_tree; prefix = "bpm:"; break;
|
||||
case DURATION_VALUE: tree = &duration_tree; prefix = "duration:"; break;
|
||||
case KEYSCALE_VALUE: tree = &keyscale_tree; prefix = "keyscale:"; break;
|
||||
case LANGUAGE_VALUE: tree = &language_tree; prefix = "language:"; break;
|
||||
case TIMESIG_VALUE: tree = ×ig_tree; prefix = "timesignature:"; break;
|
||||
default: return;
|
||||
}
|
||||
*tree = PrefixTree();
|
||||
build_value_tree(bpe, *tree, prefix, {value});
|
||||
}
|
||||
|
||||
const std::vector<int> * current_name_tokens() const {
|
||||
|
|
@ -787,11 +798,15 @@ static void parse_phase1_into_aces(
|
|||
if (!parse_cot_and_lyrics(texts[i], &parsed))
|
||||
fprintf(stderr, "WARNING: batch %d CoT parse incomplete\n", i);
|
||||
aces[i] = base;
|
||||
if (parsed.bpm > 0) aces[i].bpm = parsed.bpm;
|
||||
if (parsed.duration > 0) aces[i].duration = parsed.duration;
|
||||
if (!parsed.keyscale.empty()) aces[i].keyscale = parsed.keyscale;
|
||||
if (!parsed.timesignature.empty()) aces[i].timesignature = parsed.timesignature;
|
||||
if (!parsed.vocal_language.empty()) aces[i].vocal_language = parsed.vocal_language;
|
||||
// User-supplied metadata is authoritative. Phase 1 only fills fields
|
||||
// that were absent from the request.
|
||||
if (parsed.bpm > 0 && base.bpm <= 0) aces[i].bpm = parsed.bpm;
|
||||
if (parsed.duration > 0 && base.duration <= 0) aces[i].duration = parsed.duration;
|
||||
if (!parsed.keyscale.empty() && base.keyscale.empty()) aces[i].keyscale = parsed.keyscale;
|
||||
if (!parsed.timesignature.empty() && base.timesignature.empty()) aces[i].timesignature = parsed.timesignature;
|
||||
if (!parsed.vocal_language.empty() &&
|
||||
(base.vocal_language.empty() || base.vocal_language == "unknown"))
|
||||
aces[i].vocal_language = parsed.vocal_language;
|
||||
if (!parsed.caption.empty()) aces[i].caption = parsed.caption;
|
||||
if (merge_lyrics && !parsed.lyrics.empty()) aces[i].lyrics = parsed.lyrics;
|
||||
if (aces[i].duration <= 0) aces[i].duration = 120.0f;
|
||||
|
|
@ -1391,6 +1406,24 @@ std::string acestep_prepare_request(const music_generation_inputs inputs)
|
|||
ace.timesignature = req.timesignature;
|
||||
ace.vocal_language = req.vocal_language;
|
||||
|
||||
auto force_known_metadata = [&](MetadataFSM & target, const AcePrompt & values) {
|
||||
if (values.bpm > 0)
|
||||
target.force_field(acestep_bpe, MetadataFSM::BPM_VALUE,
|
||||
std::to_string(values.bpm));
|
||||
if (values.duration > 0)
|
||||
target.force_field(acestep_bpe, MetadataFSM::DURATION_VALUE,
|
||||
std::to_string((int)values.duration));
|
||||
if (!values.keyscale.empty())
|
||||
target.force_field(acestep_bpe, MetadataFSM::KEYSCALE_VALUE,
|
||||
values.keyscale);
|
||||
if (!values.vocal_language.empty() && values.vocal_language != "unknown")
|
||||
target.force_field(acestep_bpe, MetadataFSM::LANGUAGE_VALUE,
|
||||
values.vocal_language);
|
||||
if (!values.timesignature.empty())
|
||||
target.force_field(acestep_bpe, MetadataFSM::TIMESIG_VALUE,
|
||||
values.timesignature);
|
||||
};
|
||||
|
||||
bool user_has_codes = !req.audio_codes.empty();
|
||||
bool need_lm_codes = inputs.gen_codes && !user_has_codes;
|
||||
|
||||
|
|
@ -1421,10 +1454,10 @@ std::string acestep_prepare_request(const music_generation_inputs inputs)
|
|||
+ std::string(req.instrumental ? "true" : "false");
|
||||
prompt = build_custom_prompt(acestep_bpe, sys, user_msg.c_str());
|
||||
|
||||
// FSM: reset then optionally force language (shared for both paths)
|
||||
// Inject supplied metadata into the generated CoT so lyrics are planned
|
||||
// for the requested duration/tempo/key rather than newly guessed values.
|
||||
fsm.reset();
|
||||
if (use_fsm && ace.vocal_language != "unknown" && !ace.vocal_language.empty())
|
||||
fsm.force_language(acestep_bpe, ace.vocal_language);
|
||||
if (use_fsm) force_known_metadata(fsm, ace);
|
||||
|
||||
// Phase 1: N lyrics + metadata generations (always batched, N=batch_size)
|
||||
fprintf(stderr, "[Simple] %zu tokens, N=%d, seeds: %lld..%lld\n",
|
||||
|
|
@ -1446,8 +1479,7 @@ std::string acestep_prepare_request(const music_generation_inputs inputs)
|
|||
MetadataFSM caption_fsm;
|
||||
caption_fsm.init(acestep_bpe, acestep_llm.cfg.vocab_size, true);
|
||||
caption_fsm.reset();
|
||||
if (aces[0].vocal_language != "unknown" && !aces[0].vocal_language.empty())
|
||||
caption_fsm.force_language(acestep_bpe, aces[0].vocal_language);
|
||||
if (use_fsm) force_known_metadata(caption_fsm, aces[0]);
|
||||
|
||||
auto caption_texts = generate_phase1_batch(
|
||||
&acestep_llm, &acestep_bpe, prompt, 768,
|
||||
|
|
@ -1485,6 +1517,7 @@ std::string acestep_prepare_request(const music_generation_inputs inputs)
|
|||
prompt.size(), phase1cfg, batch_size, seed, seed + batch_size - 1);
|
||||
|
||||
fsm.reset();
|
||||
if (use_fsm) force_known_metadata(fsm, ace);
|
||||
auto phase1_texts = generate_phase1_batch(
|
||||
&acestep_llm, &acestep_bpe, prompt, 2048, temperature, top_p, top_k, rep_pen,
|
||||
seed, batch_size, use_fsm ? &fsm : nullptr, false,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue