Repository navigation
Expand file tree
/
Copy pathtypes.cpp
More file actions
90 lines (80 loc) · 2.6 KB
/
Copy pathtypes.cpp
File metadata and controls
90 lines (80 loc) · 2.6 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
#include "engine/models/yue2/types.h"
#include <stdexcept>
namespace engine::models::yue2 {
const char * cot_mode_name(Yue2CotMode mode) noexcept {
switch (mode) {
case Yue2CotMode::Off:
return "off";
case Yue2CotMode::Melody:
return "melody";
case Yue2CotMode::Full:
return "full";
}
return "full";
}
Yue2CotMode parse_cot_mode(const std::string & value) {
if (value == "off") {
return Yue2CotMode::Off;
}
if (value == "melody") {
return Yue2CotMode::Melody;
}
if (value == "full") {
return Yue2CotMode::Full;
}
throw std::runtime_error("yue2.cot must be one of off, melody, or full");
}
const char * cot_instruction(Yue2CotMode mode) noexcept {
switch (mode) {
case Yue2CotMode::Off:
return "Generate music with codec tokens from the given conditions.";
case Yue2CotMode::Melody:
return "Generate a melody-only ABC transcription without chord symbols, then generate music with codec tokens from the given conditions.";
case Yue2CotMode::Full:
return "Generate a chord-annotated ABC transcription, then generate music with codec tokens from the given conditions.";
}
return "Generate a chord-annotated ABC transcription, then generate music with codec tokens from the given conditions.";
}
const char * stop_after_name(Yue2StopAfter stage) noexcept {
switch (stage) {
case Yue2StopAfter::Abc:
return "abc";
case Yue2StopAfter::Semantic:
return "semantic";
case Yue2StopAfter::Audio:
return "audio";
}
return "audio";
}
Yue2StopAfter parse_stop_after(const std::string & value) {
if (value == "abc") {
return Yue2StopAfter::Abc;
}
if (value == "semantic") {
return Yue2StopAfter::Semantic;
}
if (value == "audio") {
return Yue2StopAfter::Audio;
}
throw std::runtime_error("yue2.stop_after must be one of abc, semantic, or audio");
}
float request_guidance_scale(const Yue2Request & request) noexcept {
if (request.cfg_scale >= 0.0F) {
return request.cfg_scale;
}
return request.cot == Yue2CotMode::Off ? 1.01F : 1.0F;
}
std::string semantic_codes_to_json(const std::vector<int32_t> & codes) {
std::string out;
out.reserve(codes.size() * 6 + 2);
out.push_back('[');
for (size_t i = 0; i < codes.size(); ++i) {
if (i != 0) {
out.push_back(',');
}
out += std::to_string(codes[i]);
}
out.push_back(']');
return out;
}
} // namespace engine::models::yue2