change custom sampler to default

This commit is contained in:
litagin02
2024-05-29 05:46:31 +09:00
parent ec12381acf
commit f4f94d4c20
4 changed files with 14 additions and 9 deletions

View File

@@ -106,7 +106,7 @@ Style-Bert-VITS2の学習用データセットを作成するためのツール
## 注意 ## 注意
- ~~長すぎる秒数12-15秒くらいより長いのwavファイルは学習に用いられないようです。また短すぎてもあまりよくない可能性もあります。~~ この制限はVer 2.5でなくなりましたが、長すぎる音声があるとVRAM消費量が増えたりするので、適度な長さにスライスすることをおすすめします。 - ~~長すぎる秒数12-15秒くらいより長いのwavファイルは学習に用いられないようです。また短すぎてもあまりよくない可能性もあります。~~ この制限はVer 2.5では学習時に「カスタムバッチサンプラーを使わない」を選択すればなくなりましたが、長すぎる音声があるとVRAM消費量が増えたり安定しなかったりするので、適度な長さにスライスすることをおすすめします。
- 書き起こしの結果をどれだけ修正すればいいかはデータセットに依存しそうです。 - 書き起こしの結果をどれだけ修正すればいいかはデータセットに依存しそうです。
""" """

View File

@@ -352,8 +352,8 @@ def train(
cmd.append("--skip_default_style") cmd.append("--skip_default_style")
if speedup: if speedup:
cmd.append("--speedup") cmd.append("--speedup")
if not not_use_custom_batch_sampler: if not_use_custom_batch_sampler:
cmd.append("--use_custom_batch_sampler") cmd.append("--not_use_custom_batch_sampler")
success, message = run_script_with_log(cmd, ignore_warning=True) success, message = run_script_with_log(cmd, ignore_warning=True)
if not success: if not success:
logger.error("Train failed.") logger.error("Train failed.")
@@ -409,9 +409,9 @@ change_log_md = """
**Ver 2.5以降の変更点** **Ver 2.5以降の変更点**
- `raw/`フォルダの中で音声をサブディレクトリに分けて配置することで、自動的にスタイルが作成されるようになりました。詳細は下の「使い方/データの前準備」を参照してください。 - `raw/`フォルダの中で音声をサブディレクトリに分けて配置することで、自動的にスタイルが作成されるようになりました。詳細は下の「使い方/データの前準備」を参照してください。
- これまでは1ファイルあたり14秒程度を超えた音声ファイルは学習には用いられていませんでしたが、Ver 2.5以降ではその制限がなくなりました。ただし: - これまでは1ファイルあたり14秒程度を超えた音声ファイルは学習には用いられていませんでしたが、Ver 2.5以降では「カスタムバッチサンプラーを使わない」にチェックを入れることでその制限が無しに学習できるようになりました(デフォルトはオフ)。ただし:
- 音声ファイルが長い場合の学習効率は悪いかもしれず、挙動も確認していません - 音声ファイルが長い場合の学習効率は悪いかもしれず、挙動も確認していません
- この変更で要求VRAMが増えるので、学習に失敗したりVRAM不足になる場合は、バッチサイズを小さくするか、学習ボタンの横の「カスタムバッチサンプラーを使う」を試してみてください(この場合は以前と同じ挙動となります)。 - この変更で要求VRAMがかなり増えるので、学習に失敗したりVRAM不足になる場合は、バッチサイズを小さくするか、チェックを外してください
""" """
how_to_md = """ how_to_md = """

View File

@@ -97,6 +97,11 @@ def run():
help="Huggingface model repo id to backup the model.", help="Huggingface model repo id to backup the model.",
default=None, default=None,
) )
parser.add_argument(
"--not_use_custom_batch_sampler",
help="Don't use custom batch sampler for training, which was used in the version < 2.5",
action="store_true",
)
args = parser.parse_args() args = parser.parse_args()
# Set log file # Set log file
@@ -213,7 +218,7 @@ def run():
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
collate_fn = TextAudioSpeakerCollate() collate_fn = TextAudioSpeakerCollate()
if args.use_custom_batch_sampler: if not args.not_use_custom_batch_sampler:
train_sampler = DistributedBucketSampler( train_sampler = DistributedBucketSampler(
train_dataset, train_dataset,
hps.train.batch_size, hps.train.batch_size,

View File

@@ -99,8 +99,8 @@ def run():
default=None, default=None,
) )
parser.add_argument( parser.add_argument(
"--use_custom_batch_sampler", "--not_use_custom_batch_sampler",
help="Use custom batch sampler for training, which was used in the version < 2.5", help="Don't use custom batch sampler for training, which was used in the version < 2.5",
action="store_true", action="store_true",
) )
args = parser.parse_args() args = parser.parse_args()
@@ -219,7 +219,7 @@ def run():
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True) collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
if args.use_custom_batch_sampler: if not args.not_use_custom_batch_sampler:
train_sampler = DistributedBucketSampler( train_sampler = DistributedBucketSampler(
train_dataset, train_dataset,
hps.train.batch_size, hps.train.batch_size,