midi-pitch-eval / score_eval.py
zziC7's picture
适配HuggingFace环境
d1abbe2
Raw
History Blame Contribute Delete
10.5 kB
import os
import numpy as np
import matplotlib
# 关键修复:防止界面环境下 Matplotlib 尝试开启 GUI 导致卡死
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import pretty_midi
from scipy.spatial.distance import cdist
class MIDIPitchEvaluator:
def __init__(self, original_midi_path, cover_midi_path,
time_threshold=0.2,
tolerance=0, unmatched_penalty_weight=0.05,
time_deviation_weight=0.01,
output_dir='results'):
"""初始化MIDI音准评价器 (Hugging Face 兼容版)"""
self.original_midi_path = original_midi_path
self.cover_midi_path = cover_midi_path
self.time_threshold = time_threshold
self.tolerance = tolerance
self.minor_deviation_upper = 1
self.severe_deviation_lower = 2
self.unmatched_penalty_weight = unmatched_penalty_weight
self.time_deviation_weight = time_deviation_weight
# 确保输出目录存在
self.output_dir = output_dir
os.makedirs(self.output_dir, exist_ok=True)
# 存储分析结果
self.original_notes = None
self.cover_notes = None
self.note_mapping = None
self.pitch_deviations = None
self.start_time_deviations = None
self.end_time_deviations = None
self.total_time_deviations = None
self.statistics = None
self.matched_pairs = []
self.unmatched_cover_notes = 0
self.unmatched_original_notes = 0
def load_midi_files(self):
try:
self.original_midi = pretty_midi.PrettyMIDI(self.original_midi_path)
self.original_notes = self._find_singing_voice_notes(self.original_midi)
self.cover_midi = pretty_midi.PrettyMIDI(self.cover_midi_path)
self.cover_notes = self._find_singing_voice_notes(self.cover_midi)
if self.original_notes is None:
if len(self.original_midi.instruments) > 0:
self.original_notes = self.original_midi.instruments[0].notes
else: return False
if self.cover_notes is None:
if len(self.cover_midi.instruments) > 0:
self.cover_notes = self.cover_midi.instruments[0].notes
else: return False
return True
except Exception as e:
print(f"加载MIDI文件出错: {e}")
return False
def _find_singing_voice_notes(self, midi_obj):
for instrument in midi_obj.instruments:
if instrument.name.lower() == "singing voice":
return instrument.notes
return None
def align_notes(self):
if not self.original_notes or not self.cover_notes:
return False
self.note_mapping = {}
self.start_time_deviations = []
self.end_time_deviations = []
self.total_time_deviations = []
self.matched_pairs = []
cover_starts = np.array([note.start for note in self.cover_notes])
cover_ends = np.array([note.end for note in self.cover_notes])
original_starts = np.array([note.start for note in self.original_notes])
original_ends = np.array([note.end for note in self.original_notes])
cover_pitches = np.array([note.pitch for note in self.cover_notes])
original_pitches = np.array([note.pitch for note in self.original_notes])
cover_time_features = np.column_stack((cover_starts, cover_ends))
original_time_features = np.column_stack((original_starts, original_ends))
time_dist_matrix = cdist(cover_time_features, original_time_features, metric='euclidean')
pitch_dist_matrix = cdist(cover_pitches.reshape(-1, 1), original_pitches.reshape(-1, 1), metric='euclidean')
for i in range(len(self.cover_notes)):
candidates = np.where(time_dist_matrix[i] < self.time_threshold * np.sqrt(2))[0]
if len(candidates) == 0:
self.note_mapping[i] = None
continue
best_j = candidates[np.argmin(pitch_dist_matrix[i, candidates])]
start_time_dev = self.cover_notes[i].start - self.original_notes[best_j].start
end_time_dev = self.cover_notes[i].end - self.original_notes[best_j].end
self.start_time_deviations.append(start_time_dev)
self.end_time_deviations.append(end_time_dev)
self.total_time_deviations.append(abs(start_time_dev) + abs(end_time_dev))
self.matched_pairs.append([i, best_j, 0, start_time_dev, end_time_dev])
self.note_mapping[i] = best_j
self.matched_count = len(self.matched_pairs)
self.total_cover_notes = len(self.cover_notes)
self.unmatched_cover_notes = self.total_cover_notes - self.matched_count
return True
def calculate_pitch_deviations(self):
if self.note_mapping is None: return False
self.pitch_deviations = []
for i in range(len(self.matched_pairs)):
cover_idx, original_idx = self.matched_pairs[i][0], self.matched_pairs[i][1]
pitch_dev = self.cover_notes[cover_idx].pitch - self.original_notes[original_idx].pitch
self.pitch_deviations.append(pitch_dev)
self.matched_pairs[i][2] = pitch_dev
return True
def calculate_statistics(self):
if not self.pitch_deviations: return False
pitch_devs = np.array(self.pitch_deviations)
total_matched = len(pitch_devs)
minor_count = np.sum((np.abs(pitch_devs) > self.tolerance) & (np.abs(pitch_devs) <= self.minor_deviation_upper))
severe_count = np.sum(np.abs(pitch_devs) >= self.severe_deviation_lower)
minor_ratio = minor_count / total_matched * 100
severe_ratio = severe_count / total_matched * 100
cover_unmatched_ratio = (self.unmatched_cover_notes / self.total_cover_notes * 100)
total_time_total = np.sum(self.total_time_deviations)
normalized_time_dev = total_time_total / (self.time_threshold * 2) if self.time_threshold > 0 else 0
time_penalty = min(normalized_time_dev * self.time_deviation_weight, 15)
base_score = 100 - severe_ratio * 0.1 - minor_ratio * 0.2
unmatched_penalty = cover_unmatched_ratio * self.unmatched_penalty_weight
total_penalty = unmatched_penalty + time_penalty
overall_score = max(60, base_score - total_penalty)
self.statistics = {
'total_matched_notes': total_matched,
'total_cover_notes': self.total_cover_notes,
'mean_abs_pitch_dev': np.mean(np.abs(pitch_devs)),
'accurate_pitch_percentage': np.sum(np.abs(pitch_devs) <= self.tolerance) / total_matched * 100,
'minor_off_percentage': minor_ratio,
'severe_off_percentage': severe_ratio,
'cover_unmatched_percentage': cover_unmatched_ratio,
'total_time_total': total_time_total,
'base_score': base_score,
'unmatched_penalty': unmatched_penalty,
'time_penalty': time_penalty,
'total_penalty': total_penalty,
'overall_score': overall_score
}
return True
def _convert_time_to_srt_format(self, seconds):
hours, rem = divmod(seconds, 3600)
minutes, seconds = divmod(rem, 60)
return f"{int(hours):02d}:{int(minutes):02d}:{int(seconds):02d},{int((seconds%1)*1000):03d}"
def generate_srt_subtitle(self):
if not self.matched_pairs: return False
sorted_pairs = sorted(self.matched_pairs, key=lambda x: self.cover_notes[x[0]].start)
srt_content = []
for i, (c_idx, o_idx, p_dev, s_dev, e_dev) in enumerate(sorted_pairs, 1):
cn, on = self.cover_notes[c_idx], self.original_notes[o_idx]
p_label = "准确" if abs(p_dev) <= self.tolerance else ("轻微走音" if abs(p_dev) <= self.minor_deviation_upper else "严重走音")
text = (f"原唱: {pretty_midi.note_number_to_name(on.pitch)} | 翻唱: {pretty_midi.note_number_to_name(cn.pitch)} | "
f"音高偏差: {p_dev:+}半音 ({p_label}) | 起始: {s_dev:+.3f}s | 结束: {e_dev:+.3f}s")
srt_content.append(f"{i}\n{self._convert_time_to_srt_format(cn.start)} --> {self._convert_time_to_srt_format(cn.end)}\n{text}\n")
out_file = os.path.join(self.output_dir, '音高与节奏对比字幕.srt')
with open(out_file, 'w', encoding='utf-8') as f:
f.write('\n'.join(srt_content))
return out_file
def visualize_results(self):
if self.statistics is None: return False
try:
plt.figure(figsize=(12, 6))
time_pts = [self.cover_notes[p[0]].start for p in self.matched_pairs]
p_devs = [p[2] for p in self.matched_pairs]
colors = ['green' if abs(d) <= self.tolerance else ('orange' if abs(d) <= self.minor_deviation_upper else 'red') for d in p_devs]
plt.scatter(time_pts, p_devs, s=15, alpha=0.7, c=colors)
plt.axhline(y=0, color='blue', linestyle='-', alpha=0.3)
plt.title('Pitch Deviation Timeline')
img_path = os.path.join(self.output_dir, '音高偏差时间序列图.png')
plt.savefig(img_path)
plt.close('all')
return img_path
except Exception as e:
print(f"可视化失败: {e}")
return None
def evaluate(self, run_visuals=True):
if not self.load_midi_files(): return False
if not self.align_notes(): return False
if not self.calculate_pitch_deviations(): return False
if not self.calculate_statistics(): return False
srt_file = self.generate_srt_subtitle()
img_file = None
if run_visuals:
img_file = self.visualize_results()
return {
"statistics": self.statistics,
"srt_file": srt_file,
"plot_img": img_file
}
if __name__ == "__main__":
# 本地测试代码保持不变,但 evaluate 返回字典
evaluator = MIDIPitchEvaluator(
"midi_input/学不会-林俊杰-align_Vocals_basic_pitch.mid",
"midi_input/学不会-zzc-align_Vocals_basic_pitch.mid"
)
result = evaluator.evaluate(run_visuals=True)
if result:
print(f"最终得分: {result['statistics']['overall_score']:.1f}")