Improve finding step count

This commit is contained in:
litagin02
2024-01-06 10:38:00 +09:00
parent 83e96a8c47
commit b31d4e7ea9

View File

@@ -1,13 +1,14 @@
import argparse import argparse
import csv import csv
import re
import warnings import warnings
from pathlib import Path from pathlib import Path
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np
import pandas as pd import pandas as pd
import torch import torch
from tqdm import tqdm from tqdm import tqdm
import numpy as np
from common.log import logger from common.log import logger
from common.tts_model import Model from common.tts_model import Model
@@ -65,6 +66,13 @@ safetensors_files = list(safetensors_files)
logger.info(f"There are {len(safetensors_files)} models.") logger.info(f"There are {len(safetensors_files)} models.")
for model_file in tqdm(safetensors_files): for model_file in tqdm(safetensors_files):
# `test_e10_s1000.safetensors`` -> 1000を取り出す
match = re.search(r"_s(\d+)\.safetensors$", model_file.name)
if match:
step = int(match.group(1))
else:
logger.warning(f"Step count not found in {model_file.name}, so skip it.")
continue
model = get_model(model_file) model = get_model(model_file)
scores = [] scores = []
for i, text in enumerate(test_texts): for i, text in enumerate(test_texts):
@@ -73,8 +81,6 @@ for model_file in tqdm(safetensors_files):
score = predictor(torch.from_numpy(audio).unsqueeze(0), sr).item() score = predictor(torch.from_numpy(audio).unsqueeze(0), sr).item()
scores.append(score) scores.append(score)
logger.info(f"score: {score}") logger.info(f"score: {score}")
# `test_e10_s1000.safetensors`` -> 1000を取り出す
step = int(model_file.stem.split("_")[2][1:])
results.append((model_file.name, step, scores)) results.append((model_file.name, step, scores))
del model del model