Fix tab_dataset

This commit is contained in:
zin
2026-05-06 23:36:28 +00:00
parent 9954603043
commit b99ca8022a
5 changed files with 307 additions and 102 deletions
+45 -46
View File
@@ -5,64 +5,63 @@ import joblib
class MusicMatcher:
def __init__(self, db_path: Path | str, model_path: Path | str):
"""
Инициализация движка сопоставления музыки.
"""
# Загружаем твою новую, обогащенную базу
self.music_db = pd.read_csv(db_path)
self.music_db['valence'] = pd.to_numeric(self.music_db['valence'], errors='coerce')
self.music_db['arousal'] = pd.to_numeric(self.music_db['arousal'], errors='coerce')
self.music_db = self.music_db.dropna()
self.acoustic_features = ['energy', 'flux', 'centroid', 'pitch', 'hnr', 'zcr']
# Удаляем строки, где нет акустических фич
self.music_db = self.music_db.dropna(subset=['valence', 'arousal'] + self.acoustic_features)
# Нормализуем акустику от 0 до 1, чтобы сравнивать с ответом LLM
self.norm_db = self.music_db.copy()
for feat in self.acoustic_features:
f_min, f_max = self.norm_db[feat].min(), self.norm_db[feat].max()
if f_max > f_min:
self.norm_db[f"norm_{feat}"] = (self.norm_db[feat] - f_min) / (f_max - f_min)
else:
self.norm_db[f"norm_{feat}"] = 0.0
self.audio_dir = Path(db_path).parent / "DEAM_audio" / "MEMD_audio"
if self.audio_dir.exists():
print(f"✅ Музыкальный архив найден: {self.audio_dir}")
else:
print(f"⚠️ ПРЕДУПРЕЖДЕНИЕ: Папка {self.audio_dir} не найдена!")
if Path(model_path).exists():
self.regressor = joblib.load(model_path)
print("✅ ML-регрессор загружен.")
else:
self.regressor = None
print("⚠️ Файл модели .pkl не найден.")
self.regressor = joblib.load(model_path) if Path(model_path).exists() else None
def predict_va(self, embedding: np.ndarray):
"""Честный прогноз координат Valence-Arousal."""
if self.regressor is not None:
emb_2d = embedding.reshape(1, -1)
prediction = self.regressor.predict(emb_2d)[0]
if self.regressor:
prediction = self.regressor.predict(embedding.reshape(1, -1))[0]
return np.clip(prediction[0], 1.0, 9.0), np.clip(prediction[1], 1.0, 9.0)
return 5.0, 5.0
def get_audio_path(self, song_id):
"""Поиск mp3 файла по его номеру."""
if not self.audio_dir.exists():
return None
if not self.audio_dir.exists(): return None
clean_id = str(int(float(song_id)))
for ext in ['.mp3', '.wav']:
file_path = self.audio_dir / f"{clean_id}{ext}"
if file_path.exists():
return file_path
path = self.audio_dir / f"{clean_id}{ext}"
if path.exists(): return path
return None
def find_nearest_tracks(self, target_v: float, target_a: float, top_k: int = 5):
"""
Поиск с использованием Взвешенного Евклидова расстояния (Weighted KNN).
Энергия (Arousal) получает больший вес, так как она сильнее
определяет жанр и ритм композиции.
"""
# Вес для Arousal = 2.0, для Valence = 1.0
# Это не позволит спокойным трекам (A < 4) попадать в выдачу
# для энергичных запросов (A > 6).
distances = np.sqrt(
1.0 * (self.music_db['valence'] - target_v)**2 +
2.5 * (self.music_db['arousal'] - target_a)**2 # Жесткий штраф за разницу в энергии
def find_nearest_tracks(self, target_v: float, target_a: float, llm_profile: dict = None, top_k: int = 5):
# 1. Эмоциональная дистанция (как и раньше)
emo_dist = np.sqrt(
1.0 * (self.norm_db['valence'] - target_v)**2 +
2.5 * (self.norm_db['arousal'] - target_a)**2
)
self.norm_db['emo_distance'] = emo_dist
df_result = self.music_db.copy()
df_result['distance'] = distances
# Сортируем по расстоянию и берем топ-K
return df_result.sort_values(by='distance').head(top_k)
# Если LLM не дала ответ, сортируем только по эмоциям
if not llm_profile:
self.norm_db['final_score'] = self.norm_db['emo_distance']
return self.norm_db.sort_values(by='final_score').head(top_k)
# 2. Акустическая дистанция (сравниваем треки с запросом LLM)
acoustic_penalty = np.zeros(len(self.norm_db))
for feat in self.acoustic_features:
if feat in llm_profile:
target_val = llm_profile[feat]
acoustic_penalty += np.abs(self.norm_db[f"norm_{feat}"] - target_val)
# Усредняем штраф
self.norm_db['acoustic_distance'] = acoustic_penalty / len(self.acoustic_features)
# 3. Финальный Score (Смесь Эмоций и Акустики). Коэф 4.0 делает акустику важной!
self.norm_db['final_score'] = self.norm_db['emo_distance'] + (self.norm_db['acoustic_distance'] * 4.0)
return self.norm_db.sort_values(by='final_score').head(top_k)