import os import sys import torch import json from trainer import Trainer, TrainerArgs # <--- КЛЮЧОВАТА ПРОМЯНА ТУК import warnings warnings.filterwarnings("ignore", category=UserWarning) warnings.filterwarnings("ignore", category=FutureWarning) # from TTS.utils.training import Trainer, TrainerArgs # from TTS.tts.utils.training import TrainerArgs from TTS.tts.configs.shared_configs import BaseDatasetConfig from TTS.tts.configs.vits_config import VitsConfig from TTS.tts.datasets import load_tts_samples from TTS.tts.models.vits import Vits, VitsArgs, VitsAudioConfig, CharactersConfig from TTS.tts.utils.speakers import SpeakerManager from TTS.tts.utils.text.tokenizer import TTSTokenizer from TTS.utils.audio import AudioProcessor from TTS.utils.manage import ModelManager # --- КОНФИГУРАЦИЯ --- MODEL_NAME = "bul-model" OUTPUT_PATH = "./my_fine_tuned_model" META_FILE_TRAIN = "Новият-Датасет2.txt" DATASET_PATH = "./dataset3" # -------------------- model_path = "./Bul-Model/model.pth" config_path = "./Bul-Model/config.json" # Изтегляне на модела # manager = ModelManager() # print(f"Изтегляне на модел: {MODEL_NAME}...") # model_path, config_path, model_item = manager.download_model(MODEL_NAME) # print(f"Моделът е изтеглен на: {model_path}") # print(f"ЗАРЕЖДАНЕ НА ОРИГИНАЛНАТА КОНФИГУРАЦИЯ") # --- ЗАРЕЖДАНЕ НА ОРИГИНАЛНАТА КОНФИГУРАЦИЯ --- with open(config_path, "r", encoding="utf-8") as f: config_dict = json.load(f) config = VitsConfig.new_from_dict(config_dict) # === КОРЕКЦИЯ НА РЕЧНИКА ЗА БЪЛГАРСКА КИРИЛИЦА === # 1. Задаваме чиста кирилица (37-те символа, включително препинателните знаци и интервала) config = VitsConfig( text_cleaner = "transliteration_cleaners", characters=CharactersConfig( characters_class="TTS.tts.models.vits.VitsCharacters", pad="", eos="", bos="", blank="", characters="абвгдежзийклмнопрстуфхцчшщъьюяѐѝ", punctuations="- _̀–", phonemes=None, ) ) # 2. Изключваме сложните мултиезични чистачи, които объркват кирилицата, # и казваме на Coqui просто да подава буквите директно (no_cleaners) # config.text_cleaner = "transliteration_cleaners" # ================================================ config.run_name = "bg" config.run_description = "Мета - БГ" config.sample_rate = 16000 config.output_path = OUTPUT_PATH config.batch_size = 32 config.eval_batch_size = 16 config.epochs = 10 config.num_workers = 4, # ВАЖНО: Казва на 4 процесорни ядра едновременно да зареждат аудио файловете предсрочно config.lr = 1e-05, config.pin_memory = True # ВАЖНО: Заключва паметта в RAM за светкавичен трансфер към видеокартата config.text_cleaner = "transliteration_cleaners" config.datasets = [BaseDatasetConfig( formatter="ljspeech", meta_file_train=META_FILE_TRAIN, language="bg", path=DATASET_PATH )] # -------------------------------------------------------- ap = AudioProcessor.init_from_config(config) tokenizer, config = TTSTokenizer.init_from_config(config) train_samples, eval_samples = load_tts_samples( config.datasets[0], eval_split=True, eval_split_max_size=config.eval_split_max_size, eval_split_size=config.eval_split_size, ) speaker_manager = SpeakerManager() speaker_manager.set_ids_from_data(train_samples + eval_samples, parse_key="speaker_name") config.model_args.num_speakers = speaker_manager.num_speakers model = Vits(config, ap, tokenizer, speaker_manager) # --- Инициализиране на трейнъра с пътя за възстановяване --- restore_path = str(model_path) CHECKPOINT_FOLDER = os.path.join(OUTPUT_PATH, 'bg', 'checkpoint') if os.path.exists(CHECKPOINT_FOLDER): all_checkpoints = sorted( [os.path.join(CHECKPOINT_FOLDER, f) for f in os.listdir(CHECKPOINT_FOLDER) if f.startswith('checkpoint')], key=os.path.getmtime, reverse=True ) if all_checkpoints: latest_checkpoint = all_checkpoints[0] print(f"Намерена е дообучена контролна точка: {latest_checkpoint}. Ще се продължи обучението от нея.") restore_path = latest_checkpoint else: print("Няма дообучена контролна точка в папката. Ще се стартира дообучение от изтегления модел.") else: print("Не е намерена папка с дообучени контролни точки. Ще се стартира дообучение от изтегления модел.") # --------------------------------------- restore_path = "" flat_checkpoint = torch.load(model_path, map_location="cpu") # Зареждаме теглата със strict=False. Това е КРИТИЧНО, защото твоят модел съдържа # тегла за дискриминатори, които чистата Vits архитектура в Coqui ще пропусне безопасно! model.load_state_dict(flat_checkpoint, strict=False) print("Моделът е зареден успешно в паметта!") # --------------------------------------- trainer_args = TrainerArgs( gpu=0, # torch.cuda.is_available(), restore_path=restore_path, ) trainer = Trainer( trainer_args, config, output_path=OUTPUT_PATH, model=model, train_samples=train_samples, eval_samples=eval_samples, ) if __name__ == '__main__': trainer.fit()