2025-11-19 01:18:16 +09:00
package main
import (
"flag"
"fmt"
"os"
"path/filepath"
"strings"
ort "github.com/yalue/onnxruntime_go"
)
// Args holds command line arguments
type Args struct {
useGPU bool
onnxDir string
totalStep int
2025-11-19 19:42:24 +09:00
speed float64
2025-11-19 01:18:16 +09:00
nTest int
voiceStyle [ ] string
text [ ] string
2026-01-06 17:15:20 +09:00
lang [ ] string
2025-11-19 01:18:16 +09:00
saveDir string
2025-11-19 18:08:30 +09:00
batch bool
2025-11-19 01:18:16 +09:00
}
func parseArgs ( ) * Args {
args := & Args { }
flag . BoolVar ( & args . useGPU , "use-gpu" , false , "Use GPU for inference (default: CPU)" )
2026-05-06 23:09:06 +02:00
flag . StringVar ( & args . onnxDir , "onnx-dir" , "../assets/onnx" , "Path to ONNX model directory" )
flag . IntVar ( & args . totalStep , "total-step" , 8 , "Number of denoising steps" )
2025-11-19 19:42:24 +09:00
flag . Float64Var ( & args . speed , "speed" , 1.05 , "Speech speed factor (higher = faster)" )
2025-11-19 01:18:16 +09:00
flag . IntVar ( & args . nTest , "n-test" , 4 , "Number of times to generate" )
flag . StringVar ( & args . saveDir , "save-dir" , "results" , "Output directory" )
2025-11-19 18:08:30 +09:00
flag . BoolVar ( & args . batch , "batch" , false , "Enable batch mode (multiple text-style pairs)" )
2025-11-19 01:18:16 +09:00
2026-01-06 17:15:20 +09:00
var voiceStyleStr , textStr , langStr string
2026-05-06 23:09:06 +02:00
flag . StringVar ( & voiceStyleStr , "voice-style" , "../assets/voice_styles/M1.json" , "Voice style file path(s), comma-separated" )
2025-11-19 01:18:16 +09:00
flag . StringVar ( & textStr , "text" , "This morning, I took a walk in the park, and the sound of the birds and the breeze was so pleasant that I stopped for a long time just to listen." , "Text(s) to synthesize, pipe-separated" )
2026-05-06 23:09:06 +02:00
flag . StringVar ( & langStr , "lang" , "en" , "Language(s) for synthesis, comma-separated" )
2025-11-19 01:18:16 +09:00
flag . Parse ( )
// Parse comma-separated voice-style
if voiceStyleStr != "" {
args . voiceStyle = strings . Split ( voiceStyleStr , "," )
for i := range args . voiceStyle {
args . voiceStyle [ i ] = strings . TrimSpace ( args . voiceStyle [ i ] )
}
}
// Parse pipe-separated text
if textStr != "" {
args . text = strings . Split ( textStr , "|" )
for i := range args . text {
args . text [ i ] = strings . TrimSpace ( args . text [ i ] )
}
}
2026-01-06 17:15:20 +09:00
// Parse comma-separated lang
if langStr != "" {
args . lang = strings . Split ( langStr , "," )
for i := range args . lang {
args . lang [ i ] = strings . TrimSpace ( args . lang [ i ] )
}
}
2025-11-19 01:18:16 +09:00
return args
}
func main ( ) {
fmt . Println ( "=== TTS Inference with ONNX Runtime (Go) ===\n" )
// --- 1. Parse arguments --- //
args := parseArgs ( )
totalStep := args . totalStep
2025-11-19 19:42:24 +09:00
speed := float32 ( args . speed )
2025-11-19 01:18:16 +09:00
nTest := args . nTest
saveDir := args . saveDir
voiceStylePaths := args . voiceStyle
textList := args . text
2026-01-06 17:15:20 +09:00
langList := args . lang
2025-11-19 18:08:30 +09:00
batch := args . batch
2025-11-19 01:18:16 +09:00
2025-11-19 18:08:30 +09:00
if batch {
if len ( voiceStylePaths ) != len ( textList ) {
fmt . Printf ( "Error: Number of voice styles (%d) must match number of texts (%d)\n" ,
len ( voiceStylePaths ) , len ( textList ) )
os . Exit ( 1 )
}
2026-01-06 17:15:20 +09:00
if len ( langList ) != len ( textList ) {
fmt . Printf ( "Error: Number of languages (%d) must match number of texts (%d)\n" ,
len ( langList ) , len ( textList ) )
os . Exit ( 1 )
}
2025-11-19 01:18:16 +09:00
}
bsz := len ( voiceStylePaths )
// Initialize ONNX Runtime
if err := InitializeONNXRuntime ( ) ; err != nil {
fmt . Printf ( "Error initializing ONNX Runtime: %v\n" , err )
os . Exit ( 1 )
}
defer ort . DestroyEnvironment ( )
// --- 2. Load config --- //
cfg , err := LoadCfgs ( args . onnxDir )
if err != nil {
fmt . Printf ( "Error loading config: %v\n" , err )
os . Exit ( 1 )
}
// --- 3. Load TTS components --- //
textToSpeech , err := LoadTextToSpeech ( args . onnxDir , args . useGPU , cfg )
if err != nil {
fmt . Printf ( "Error loading TTS components: %v\n" , err )
os . Exit ( 1 )
}
defer textToSpeech . Destroy ( )
// --- 4. Load voice styles --- //
style , err := LoadVoiceStyle ( voiceStylePaths , true )
if err != nil {
fmt . Printf ( "Error loading voice styles: %v\n" , err )
os . Exit ( 1 )
}
defer style . Destroy ( )
// --- 5. Synthesize speech --- //
if err := os . MkdirAll ( saveDir , 0755 ) ; err != nil {
fmt . Printf ( "Error creating save directory: %v\n" , err )
os . Exit ( 1 )
}
for n := 0 ; n < nTest ; n ++ {
fmt . Printf ( "\n[%d/%d] Starting synthesis...\n" , n + 1 , nTest )
var wav [ ] float32
var duration [ ] float32
2025-11-19 18:08:30 +09:00
if batch {
Timer ( "Generating speech from text" , func ( ) interface { } {
2026-01-06 17:15:20 +09:00
w , d , err := textToSpeech . Batch ( textList , langList , style , totalStep , speed )
2025-11-19 18:08:30 +09:00
if err != nil {
fmt . Printf ( "Error generating speech: %v\n" , err )
os . Exit ( 1 )
}
wav = w
duration = d
return nil
} )
} else {
Timer ( "Generating speech from text" , func ( ) interface { } {
2026-01-06 17:15:20 +09:00
w , d , err := textToSpeech . Call ( textList [ 0 ] , langList [ 0 ] , style , totalStep , speed , 0.3 )
2025-11-19 18:08:30 +09:00
if err != nil {
fmt . Printf ( "Error generating speech: %v\n" , err )
os . Exit ( 1 )
}
wav = w
duration = [ ] float32 { d }
return nil
} )
}
2025-11-19 01:18:16 +09:00
// Save outputs
for i := 0 ; i < bsz ; i ++ {
fname := fmt . Sprintf ( "%s_%d.wav" , sanitizeFilename ( textList [ i ] , 20 ) , n + 1 )
2025-11-19 18:08:30 +09:00
var wavOut [ ] float64
if batch {
wavOut = extractWavSegment ( wav , duration [ i ] , textToSpeech . SampleRate , i , bsz )
} else {
// For non-batch mode, wav is a single concatenated audio
wavLen := int ( float32 ( textToSpeech . SampleRate ) * duration [ 0 ] )
wavOut = make ( [ ] float64 , wavLen )
for j := 0 ; j < wavLen && j < len ( wav ) ; j ++ {
wavOut [ j ] = float64 ( wav [ j ] )
}
}
2025-11-19 01:18:16 +09:00
outputPath := filepath . Join ( saveDir , fname )
if err := writeWavFile ( outputPath , wavOut , textToSpeech . SampleRate ) ; err != nil {
fmt . Printf ( "Error writing wav file: %v\n" , err )
continue
}
fmt . Printf ( "Saved: %s\n" , outputPath )
}
}
fmt . Println ( "\n=== Synthesis completed successfully! ===" )
}