File size: 5,111 Bytes
ba07985
 
 
 
 
 
 
 
3e08188
ba07985
 
 
 
 
e1ec36f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ba07985
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3e08188
e1ec36f
 
 
 
 
 
 
ba07985
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#include "ax_engine.hpp"
#include "kantts.hpp"

#include <chrono>
#include <cmath>
#include <cstdio>
#include <fstream>
#include <iostream>
#include <sstream>
#include <string>
#include <vector>

using namespace kantts;

// duration 模型固定输入 22 帧:把超过 22 个有效符号的长句按标点切成子段,
// 每段独立走完整管线(enc/韵律/时长/解码/voc),音频按段拼接。
static const int kMaxDurT = 22;

static bool IsPunct(const std::string& tok) {
    return tok.size() >= 2 && tok[0] == '{' && tok[1] == '#';
}

static std::string JoinSymbols(const std::vector<std::string>& toks, bool with_end) {
    std::string s;
    for (const auto& t : toks) {
        if (!s.empty()) s += ' ';
        s += t;
    }
    if (with_end) {
        if (!s.empty()) s += ' ';
        s += "~";
    }
    return s;
}

static std::vector<std::string> SplitLongSentence(const std::string& line) {
    std::vector<std::string> toks;
    std::stringstream ss(line);
    std::string t;
    while (ss >> t) toks.push_back(t);
    if (toks.empty()) return {};
    // 末尾 ~ 结束符单独保留
    bool has_end = false;
    if (toks.back() == "~") {
        has_end = true;
        toks.pop_back();
    }
    // duration 模型的 T = Encode 后的符号数(含标点),上限 kMaxDurT
    int total_valid = (int)toks.size();
    if (total_valid <= kMaxDurT) {
        return {line};
    }
    // 按标点切块
    std::vector<std::vector<std::string>> blocks;
    std::vector<std::string> cur;
    for (const auto& s : toks) {
        cur.push_back(s);
        if (IsPunct(s)) {
            blocks.push_back(cur);
            cur.clear();
        }
    }
    if (!cur.empty()) blocks.push_back(cur);
    // 贪心合并块,使每段有效符号 ≤ kMaxDurT
    std::vector<std::string> out;
    std::vector<std::string> merged;
    int merged_valid = 0;
    for (const auto& blk : blocks) {
        int blk_valid = (int)blk.size();
        if (!merged.empty() && merged_valid + blk_valid > kMaxDurT) {
            out.push_back(JoinSymbols(merged, has_end));
            merged.clear();
            merged_valid = 0;
        }
        merged.insert(merged.end(), blk.begin(), blk.end());
        merged_valid += blk_valid;
    }
    if (!merged.empty()) out.push_back(JoinSymbols(merged, has_end));
    return out;
}

static void WriteWav(const std::string& path, const std::vector<float>& audio, int sr = 16000) {
    std::vector<int16_t> pcm(audio.size());
    for (size_t i = 0; i < audio.size(); ++i) {
        float v = audio[i];
        if (v > 1.0f) v = 1.0f;
        if (v < -1.0f) v = -1.0f;
        pcm[i] = (int16_t)(v * 32767.0f);
    }
    std::ofstream f(path, std::ios::binary);
    auto wr = [&](const void* p, size_t n) { f.write((const char*)p, n); };
    uint32_t data = pcm.size() * 2;
    uint32_t rate = sr;
    uint16_t ch = 1, bits = 16;
    wr("RIFF", 4);
    uint32_t riff_size = 36 + data;
    wr(&riff_size, 4);
    wr("WAVEfmt ", 8);
    uint32_t hdr = 16;
    uint16_t fmt = 1;
    wr(&hdr, 4);
    wr(&fmt, 2);
    wr(&ch, 2);
    wr(&rate, 4);
    uint32_t bps = rate * ch * bits / 8;
    wr(&bps, 4);
    uint16_t ba = ch * bits / 8;
    wr(&ba, 2);
    wr(&bits, 2);
    wr("data", 4);
    wr(&data, 4);
    wr(pcm.data(), pcm.size() * 2);
}

int main(int argc, char** argv) {
    if (argc < 4) {
        std::fprintf(stderr,
                     "用法: kantts_tts <model_dir> <resource_dir> <symbols.txt> <out.wav>\n"
                     "symbols.txt: 每行一个 ttsfrd gen_tacotron_symbols 输出(见 tools/text_to_symbols.py)\n");
        return 1;
    }
    try {
        AxRuntimeInit();
        KanttsPipeline pipe(argv[1], argv[2], std::string(argv[1]) + "/am_config.yaml");
        std::ifstream sf(argv[3]);
        std::vector<std::string> symbols;
        std::string line;
        while (std::getline(sf, line)) {
            auto tab = line.find('\t');
            std::string sym = tab == std::string::npos ? line : line.substr(tab + 1);
            if (std::getenv("KANTTS_CPU_LONG")) {
                symbols.push_back(sym);  // 整句走 CPU 长句管线(对照用)
            } else {
                auto parts = SplitLongSentence(sym);  // 默认:按标点切分,子段全走 NPU
                symbols.insert(symbols.end(), parts.begin(), parts.end());
                if (parts.size() > 1) std::fprintf(stderr, "[stage] 长句切分为 %zu 段\n", parts.size());
            }
        }
        auto t0 = std::chrono::steady_clock::now();
        auto audio = pipe.SynthesizeSymbols(symbols);
        auto t1 = std::chrono::steady_clock::now();
        WriteWav(argv[4], audio);
        double sec = std::chrono::duration<double>(t1 - t0).count();
        double dur = audio.size() / 16000.0;
        std::printf("输出 %s(%.2fs 音频,合成 %.2fs,RTF=%.2f)\n", argv[4], dur, sec,
                    dur > 0 ? sec / dur : 0);
        AxRuntimeDeinit();
    } catch (const std::exception& e) {
        std::fprintf(stderr, "错误: %s\n", e.what());
        return 1;
    }
    return 0;
}