ForyoClassifierCNN: A music classification project based on PyTorch¶
Overview¶
The project uses PyTorch to implement a CNN model for classifying atarayo, yorushika, yoasobi and zutomayo, which is suitable for beginners who ardently love Japanese music (such as me) in computer audition and music information retrieval to learn and practice.
- Repository: ForyoClassifierCNN | GitHub
This notebook is for exploring and analyzing the music dataset, including:
Dataset statistics
Audio feature visualization
Feature distribution analysis
Model training results analysis
In [1]:
Copied!
import sys
sys.path.append('../')
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
import librosa
import librosa.display
from pathlib import Path
import json
import yaml
import warnings
from src.feature_extraction import extract_features
from src.data_preparation import get_audio_files
warnings.filterwarnings('ignore')
import sys sys.path.append('../') import numpy as np import pandas as pd import matplotlib.pyplot as plt import seaborn as sns import librosa import librosa.display from pathlib import Path import json import yaml import warnings from src.feature_extraction import extract_features from src.data_preparation import get_audio_files warnings.filterwarnings('ignore')
1. Load configuration and data¶
In [2]:
Copied!
# Load configuration
with open('../config.yaml', 'r', encoding='utf-8') as f:
config = yaml.safe_load(f)
data_config = config['data']
feature_config = config['features']
print("Configuration:")
print(f" Sample rate: {data_config['sample_rate']} Hz")
print(f" Duration: {data_config.get('duration', 'None')} seconds")
print(f" MFCC dimension: {feature_config['n_mfcc']}")
print(f" Include Chroma feature: {feature_config.get('include_chroma', False)}")
print(f" Include spectral feature: {feature_config.get('include_spectral', False)}")
# Load configuration with open('../config.yaml', 'r', encoding='utf-8') as f: config = yaml.safe_load(f) data_config = config['data'] feature_config = config['features'] print("Configuration:") print(f" Sample rate: {data_config['sample_rate']} Hz") print(f" Duration: {data_config.get('duration', 'None')} seconds") print(f" MFCC dimension: {feature_config['n_mfcc']}") print(f" Include Chroma feature: {feature_config.get('include_chroma', False)}") print(f" Include spectral feature: {feature_config.get('include_spectral', False)}")
Configuration: Sample rate: 22050 Hz Duration: 30 seconds MFCC dimension: 13 Include Chroma feature: True Include spectral feature: True
In [3]:
Copied!
# Get all audio files
# Adjust path relative to notebook location (notebooks/ directory)
raw_dir = data_config['raw_dir']
if not Path(raw_dir).is_absolute():
# Convert relative path from project root to relative path from notebook
raw_dir = str(Path('../') / raw_dir)
audio_files = get_audio_files(raw_dir)
if not audio_files:
print(f"Warning: No audio files found in {raw_dir}!")
print("Please ensure the data directory structure is as follows:")
print(" data/raw/")
print(" ├── band1/")
print(" │ ├── song1.wav")
print(" └── ...")
else:
print(f"Found {len(audio_files)} bands:")
total_files = 0
for band_name, files in audio_files.items():
print(f" {band_name}: {len(files)} files")
total_files += len(files)
print(f"\nTotal: {total_files} audio files")
# Get all audio files # Adjust path relative to notebook location (notebooks/ directory) raw_dir = data_config['raw_dir'] if not Path(raw_dir).is_absolute(): # Convert relative path from project root to relative path from notebook raw_dir = str(Path('../') / raw_dir) audio_files = get_audio_files(raw_dir) if not audio_files: print(f"Warning: No audio files found in {raw_dir}!") print("Please ensure the data directory structure is as follows:") print(" data/raw/") print(" ├── band1/") print(" │ ├── song1.wav") print(" └── ...") else: print(f"Found {len(audio_files)} bands:") total_files = 0 for band_name, files in audio_files.items(): print(f" {band_name}: {len(files)} files") total_files += len(files) print(f"\nTotal: {total_files} audio files")
Found 4 bands: atarayo: 82 files yoasobi: 107 files yorushika: 237 files zutomayo: 203 files Total: 629 audio files
2. Audio feature visualization¶
In [4]:
Copied!
# Visualize the features of the first audio file
if audio_files:
band_name = list(audio_files.keys())[0]
audio_path = audio_files[band_name][0]
print(f"Analyzing audio: {Path(audio_path).name}")
print(f"Band: {band_name}")
# Load audio
audio, sr = librosa.load(audio_path, sr=data_config['sample_rate'])
duration = len(audio) / sr
print(f"Duration: {duration:.2f} seconds")
# Extract features
features = extract_features(audio, sr=sr, config=config)
print(f"\nFeature shape: {features.shape}")
print(f" - Feature dimension: {features.shape[0]}")
print(f" - Time frames: {features.shape[1]}")
# Plot MFCC features
plt.figure(figsize=(14, 6))
mfcc_features = features[:feature_config['n_mfcc'], :]
librosa.display.specshow(
mfcc_features,
x_axis='time',
sr=sr,
hop_length=feature_config['hop_length'],
cmap='viridis'
)
plt.colorbar(format='%+2.0f')
plt.title(f'MFCC Features - {band_name}', fontsize=14, fontweight='bold')
plt.ylabel('MFCC Coefficient')
plt.tight_layout()
plt.show()
else:
print("No available audio files for visualization")
# Visualize the features of the first audio file if audio_files: band_name = list(audio_files.keys())[0] audio_path = audio_files[band_name][0] print(f"Analyzing audio: {Path(audio_path).name}") print(f"Band: {band_name}") # Load audio audio, sr = librosa.load(audio_path, sr=data_config['sample_rate']) duration = len(audio) / sr print(f"Duration: {duration:.2f} seconds") # Extract features features = extract_features(audio, sr=sr, config=config) print(f"\nFeature shape: {features.shape}") print(f" - Feature dimension: {features.shape[0]}") print(f" - Time frames: {features.shape[1]}") # Plot MFCC features plt.figure(figsize=(14, 6)) mfcc_features = features[:feature_config['n_mfcc'], :] librosa.display.specshow( mfcc_features, x_axis='time', sr=sr, hop_length=feature_config['hop_length'], cmap='viridis' ) plt.colorbar(format='%+2.0f') plt.title(f'MFCC Features - {band_name}', fontsize=14, fontweight='bold') plt.ylabel('MFCC Coefficient') plt.tight_layout() plt.show() else: print("No available audio files for visualization")
In [5]:
Copied!
# Plot complete feature set (including MFCC, Chroma, etc.)
if audio_files and len(audio_files) > 0:
band_name = list(audio_files.keys())[0]
audio_path = audio_files[band_name][0]
audio, sr = librosa.load(audio_path, sr=data_config['sample_rate'])
features = extract_features(audio, sr=sr, config=config)
fig, axes = plt.subplots(features.shape[0] // 4 + 1, 1, figsize=(14, 3 * (features.shape[0] // 4 + 1)))
if features.shape[0] <= 4:
axes = [axes]
else:
axes = axes.flatten()
feature_idx = 0
feature_names = []
# MFCC feature
n_mfcc = feature_config['n_mfcc']
for i in range(min(n_mfcc, features.shape[0])):
if feature_idx < len(axes):
librosa.display.specshow(
features[i:i+1, :],
x_axis='time',
sr=sr,
hop_length=feature_config['hop_length'],
ax=axes[feature_idx],
cmap='viridis'
)
axes[feature_idx].set_title(f'MFCC-{i}')
axes[feature_idx].set_ylabel('')
feature_idx += 1
# Chroma feature
if feature_config.get('include_chroma', False) and feature_idx + 12 <= features.shape[0]:
chroma_start = n_mfcc
for i in range(12):
if feature_idx < len(axes):
librosa.display.specshow(
features[chroma_start + i:chroma_start + i + 1, :],
x_axis='time',
sr=sr,
hop_length=feature_config['hop_length'],
ax=axes[feature_idx],
cmap='viridis'
)
axes[feature_idx].set_title(f'Chroma-{i}')
axes[feature_idx].set_ylabel('')
feature_idx += 1
plt.tight_layout()
plt.show()
# Plot complete feature set (including MFCC, Chroma, etc.) if audio_files and len(audio_files) > 0: band_name = list(audio_files.keys())[0] audio_path = audio_files[band_name][0] audio, sr = librosa.load(audio_path, sr=data_config['sample_rate']) features = extract_features(audio, sr=sr, config=config) fig, axes = plt.subplots(features.shape[0] // 4 + 1, 1, figsize=(14, 3 * (features.shape[0] // 4 + 1))) if features.shape[0] <= 4: axes = [axes] else: axes = axes.flatten() feature_idx = 0 feature_names = [] # MFCC feature n_mfcc = feature_config['n_mfcc'] for i in range(min(n_mfcc, features.shape[0])): if feature_idx < len(axes): librosa.display.specshow( features[i:i+1, :], x_axis='time', sr=sr, hop_length=feature_config['hop_length'], ax=axes[feature_idx], cmap='viridis' ) axes[feature_idx].set_title(f'MFCC-{i}') axes[feature_idx].set_ylabel('') feature_idx += 1 # Chroma feature if feature_config.get('include_chroma', False) and feature_idx + 12 <= features.shape[0]: chroma_start = n_mfcc for i in range(12): if feature_idx < len(axes): librosa.display.specshow( features[chroma_start + i:chroma_start + i + 1, :], x_axis='time', sr=sr, hop_length=feature_config['hop_length'], ax=axes[feature_idx], cmap='viridis' ) axes[feature_idx].set_title(f'Chroma-{i}') axes[feature_idx].set_ylabel('') feature_idx += 1 plt.tight_layout() plt.show()
3. Feature comparison between different bands¶
In [6]:
Copied!
# Analyze the feature distribution of different bands
if len(audio_files) >= 4:
fig, axes = plt.subplots(2, 2, figsize=(16, 12))
axes = axes.flatten()
for idx, (band_name, files) in enumerate(list(audio_files.items())[:4]):
# Randomly select a song
audio_path = np.random.choice(files)
audio, sr = librosa.load(audio_path, sr=data_config['sample_rate'])
features = extract_features(audio, sr=sr, config=config)
# Plot MFCC features
ax = axes[idx]
mfcc_features = features[:feature_config['n_mfcc'], :]
librosa.display.specshow(
mfcc_features,
x_axis='time',
sr=sr,
hop_length=feature_config['hop_length'],
ax=ax,
cmap='viridis'
)
ax.set_title(f'{band_name}', fontsize=12, fontweight='bold')
ax.set_ylabel('MFCC Coefficient')
plt.tight_layout()
plt.show()
else:
print(f"Currently only {len(audio_files)} bands, at least 4 bands are needed for comparison")
# Analyze the feature distribution of different bands if len(audio_files) >= 4: fig, axes = plt.subplots(2, 2, figsize=(16, 12)) axes = axes.flatten() for idx, (band_name, files) in enumerate(list(audio_files.items())[:4]): # Randomly select a song audio_path = np.random.choice(files) audio, sr = librosa.load(audio_path, sr=data_config['sample_rate']) features = extract_features(audio, sr=sr, config=config) # Plot MFCC features ax = axes[idx] mfcc_features = features[:feature_config['n_mfcc'], :] librosa.display.specshow( mfcc_features, x_axis='time', sr=sr, hop_length=feature_config['hop_length'], ax=ax, cmap='viridis' ) ax.set_title(f'{band_name}', fontsize=12, fontweight='bold') ax.set_ylabel('MFCC Coefficient') plt.tight_layout() plt.show() else: print(f"Currently only {len(audio_files)} bands, at least 4 bands are needed for comparison")
[src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 31 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken. [src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 9 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken. [src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 9 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken. [src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 9 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken. [src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 7 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken. [src/libmpg123/id3.c:INT123_id3_to_utf8():394] warning: Weird tag size 31 for encoding 1 - I will probably trim too early or something but I think the MP3 is broken.
4. Audio duration distribution analysis¶
In [7]:
Copied!
# Analyze the audio duration distribution
if audio_files:
durations = []
bands = []
print("Analyzing audio duration...")
for band_name, files in audio_files.items():
for audio_path in files[:20]: # Limit each band to 20 songs to avoid excessive calculation time
try:
duration = librosa.get_duration(filename=audio_path)
durations.append(duration)
bands.append(band_name)
except Exception as e:
print(f"Cannot read {audio_path}: {e}")
continue
if durations:
df = pd.DataFrame({'duration': durations, 'band': bands})
# Plot boxplot
plt.figure(figsize=(12, 6))
sns.boxplot(data=df, x='band', y='duration')
plt.title('Audio Duration Distribution (by Band)', fontsize=14, fontweight='bold')
plt.xlabel('Band', fontsize=12)
plt.ylabel('Duration (seconds)', fontsize=12)
plt.xticks(rotation=45)
plt.tight_layout()
plt.show()
print("\nDuration statistics:")
print(df.groupby('band')['duration'].describe())
else:
print("No audio duration information available")
# Analyze the audio duration distribution if audio_files: durations = [] bands = [] print("Analyzing audio duration...") for band_name, files in audio_files.items(): for audio_path in files[:20]: # Limit each band to 20 songs to avoid excessive calculation time try: duration = librosa.get_duration(filename=audio_path) durations.append(duration) bands.append(band_name) except Exception as e: print(f"Cannot read {audio_path}: {e}") continue if durations: df = pd.DataFrame({'duration': durations, 'band': bands}) # Plot boxplot plt.figure(figsize=(12, 6)) sns.boxplot(data=df, x='band', y='duration') plt.title('Audio Duration Distribution (by Band)', fontsize=14, fontweight='bold') plt.xlabel('Band', fontsize=12) plt.ylabel('Duration (seconds)', fontsize=12) plt.xticks(rotation=45) plt.tight_layout() plt.show() print("\nDuration statistics:") print(df.groupby('band')['duration'].describe()) else: print("No audio duration information available")
Analyzing audio duration...
Duration statistics:
count mean std min 25% 50% \
band
atarayo 20.0 268.952074 47.372627 204.667029 230.264308 258.755510
yoasobi 20.0 220.372243 39.008887 182.702500 197.105768 203.878062
yorushika 20.0 250.779988 27.900556 200.186896 234.280663 250.578250
zutomayo 20.0 250.642994 26.325979 214.466667 233.500000 249.519943
75% max
band
atarayo 300.502834 358.356440
yoasobi 242.903350 335.465260
yorushika 268.339242 299.792834
zutomayo 259.656667 326.266667
5. MFCC feature statistics comparison¶
In [8]:
Copied!
# Compare the statistics of MFCC features
if audio_files:
band_features = {}
print("Extracting features...")
for band_name, files in audio_files.items():
mfccs_list = []
for audio_path in files[:10]: # Each band takes 10 songs
try:
audio, sr = librosa.load(audio_path, sr=data_config['sample_rate'])
features = extract_features(audio, sr=sr, config=config)
mfcc = features[:feature_config['n_mfcc'], :]
# Calculate time average
mfcc_mean = np.mean(mfcc, axis=1)
mfccs_list.append(mfcc_mean)
except Exception:
continue
if mfccs_list:
band_features[band_name] = np.array(mfccs_list)
print(f" {band_name}: {len(mfccs_list)} samples")
# Plot average MFCC features
if band_features:
fig, ax = plt.subplots(figsize=(14, 6))
for band_name, mfccs in band_features.items():
mean_mfcc = np.mean(mfccs, axis=0)
std_mfcc = np.std(mfccs, axis=0)
x = np.arange(len(mean_mfcc))
ax.plot(x, mean_mfcc, label=band_name, linewidth=2)
ax.fill_between(x, mean_mfcc - std_mfcc, mean_mfcc + std_mfcc, alpha=0.2)
ax.set_xlabel('MFCC Coefficient Index', fontsize=12)
ax.set_ylabel('Average MFCC Value', fontsize=12)
ax.set_title('Average MFCC Features by Band', fontsize=14, fontweight='bold')
ax.legend()
ax.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
# Compare the statistics of MFCC features if audio_files: band_features = {} print("Extracting features...") for band_name, files in audio_files.items(): mfccs_list = [] for audio_path in files[:10]: # Each band takes 10 songs try: audio, sr = librosa.load(audio_path, sr=data_config['sample_rate']) features = extract_features(audio, sr=sr, config=config) mfcc = features[:feature_config['n_mfcc'], :] # Calculate time average mfcc_mean = np.mean(mfcc, axis=1) mfccs_list.append(mfcc_mean) except Exception: continue if mfccs_list: band_features[band_name] = np.array(mfccs_list) print(f" {band_name}: {len(mfccs_list)} samples") # Plot average MFCC features if band_features: fig, ax = plt.subplots(figsize=(14, 6)) for band_name, mfccs in band_features.items(): mean_mfcc = np.mean(mfccs, axis=0) std_mfcc = np.std(mfccs, axis=0) x = np.arange(len(mean_mfcc)) ax.plot(x, mean_mfcc, label=band_name, linewidth=2) ax.fill_between(x, mean_mfcc - std_mfcc, mean_mfcc + std_mfcc, alpha=0.2) ax.set_xlabel('MFCC Coefficient Index', fontsize=12) ax.set_ylabel('Average MFCC Value', fontsize=12) ax.set_title('Average MFCC Features by Band', fontsize=14, fontweight='bold') ax.legend() ax.grid(True, alpha=0.3) plt.tight_layout() plt.show()
6. Dataset split statistics¶
In [9]:
Copied!
# Check the dataset split
splits_dir = Path('../data/splits')
if splits_dir.exists():
try:
with open(splits_dir / 'train_files.json', 'r') as f:
train_files = json.load(f)
with open(splits_dir / 'val_files.json', 'r') as f:
val_files = json.load(f)
with open(splits_dir / 'test_files.json', 'r') as f:
test_files = json.load(f)
# Statistics of files in each dataset
split_stats = {
'Train': {band: len(files) for band, files in train_files.items()},
'Validation': {band: len(files) for band, files in val_files.items()},
'Test': {band: len(files) for band, files in test_files.items()}
}
df_splits = pd.DataFrame(split_stats).T
print("Dataset split statistics:")
print(df_splits)
# Plot
df_splits.plot(kind='bar', figsize=(12, 6), rot=0)
plt.title('Dataset Split Statistics', fontsize=14, fontweight='bold')
plt.xlabel('Split', fontsize=12)
plt.ylabel('Number of Files', fontsize=12)
plt.legend(title='Band')
plt.tight_layout()
plt.show()
except FileNotFoundError:
print("Dataset split files not found, please run data_preparation.py first")
else:
print("splits directory not found, please run data_preparation.py first")
# Check the dataset split splits_dir = Path('../data/splits') if splits_dir.exists(): try: with open(splits_dir / 'train_files.json', 'r') as f: train_files = json.load(f) with open(splits_dir / 'val_files.json', 'r') as f: val_files = json.load(f) with open(splits_dir / 'test_files.json', 'r') as f: test_files = json.load(f) # Statistics of files in each dataset split_stats = { 'Train': {band: len(files) for band, files in train_files.items()}, 'Validation': {band: len(files) for band, files in val_files.items()}, 'Test': {band: len(files) for band, files in test_files.items()} } df_splits = pd.DataFrame(split_stats).T print("Dataset split statistics:") print(df_splits) # Plot df_splits.plot(kind='bar', figsize=(12, 6), rot=0) plt.title('Dataset Split Statistics', fontsize=14, fontweight='bold') plt.xlabel('Split', fontsize=12) plt.ylabel('Number of Files', fontsize=12) plt.legend(title='Band') plt.tight_layout() plt.show() except FileNotFoundError: print("Dataset split files not found, please run data_preparation.py first") else: print("splits directory not found, please run data_preparation.py first")
7. Model training results analysis¶
In [10]:
Copied!
# Load training history (if exists)
history_path = Path('../logs/training_history.json')
if history_path.exists():
with open(history_path, 'r') as f:
history = json.load(f)
# Plot training curves
epochs = range(1, len(history['train_loss']) + 1)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))
# Loss curve
ax1.plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2)
ax1.plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)
ax1.set_xlabel('Epoch', fontsize=12)
ax1.set_ylabel('Loss', fontsize=12)
ax1.set_title('Training and Validation Loss', fontsize=14, fontweight='bold')
ax1.legend()
ax1.grid(True, alpha=0.3)
# Accuracy curve
ax2.plot(epochs, history['train_acc'], 'b-', label='Train Acc', linewidth=2)
ax2.plot(epochs, history['val_acc'], 'r-', label='Val Acc', linewidth=2)
ax2.set_xlabel('Epoch', fontsize=12)
ax2.set_ylabel('Accuracy (%)', fontsize=12)
ax2.set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')
ax2.legend()
ax2.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
# Print final metrics
print("Final training results:")
print(f" Final train loss: {history['train_loss'][-1]:.4f}")
print(f" Final val loss: {history['val_loss'][-1]:.4f}")
print(f" Final train acc: {history['train_acc'][-1]:.2f}%")
print(f" Final val acc: {history['val_acc'][-1]:.2f}%")
else:
print("Training history file not found, please train the model first")
# Load training history (if exists) history_path = Path('../logs/training_history.json') if history_path.exists(): with open(history_path, 'r') as f: history = json.load(f) # Plot training curves epochs = range(1, len(history['train_loss']) + 1) fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5)) # Loss curve ax1.plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2) ax1.plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2) ax1.set_xlabel('Epoch', fontsize=12) ax1.set_ylabel('Loss', fontsize=12) ax1.set_title('Training and Validation Loss', fontsize=14, fontweight='bold') ax1.legend() ax1.grid(True, alpha=0.3) # Accuracy curve ax2.plot(epochs, history['train_acc'], 'b-', label='Train Acc', linewidth=2) ax2.plot(epochs, history['val_acc'], 'r-', label='Val Acc', linewidth=2) ax2.set_xlabel('Epoch', fontsize=12) ax2.set_ylabel('Accuracy (%)', fontsize=12) ax2.set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold') ax2.legend() ax2.grid(True, alpha=0.3) plt.tight_layout() plt.show() # Print final metrics print("Final training results:") print(f" Final train loss: {history['train_loss'][-1]:.4f}") print(f" Final val loss: {history['val_loss'][-1]:.4f}") print(f" Final train acc: {history['train_acc'][-1]:.2f}%") print(f" Final val acc: {history['val_acc'][-1]:.2f}%") else: print("Training history file not found, please train the model first")
In [11]:
Copied!
# Load evaluation results (if exists)
eval_path = Path('../results/evaluation_results.json')
if eval_path.exists():
with open(eval_path, 'r', encoding='utf-8') as f:
eval_results = json.load(f)
print("=" * 70)
print("Model evaluation results")
print("=" * 70)
print(f"Test set accuracy: {eval_results['accuracy'] * 100:.2f}%")
print("\nMacro average:")
print(f" Precision: {eval_results['precision_macro'] * 100:.2f}%")
print(f" Recall: {eval_results['recall_macro'] * 100:.2f}%")
print(f" F1 score: {eval_results['f1_macro'] * 100:.2f}%")
print("\nDetailed metrics for each class:")
print("-" * 70)
for class_name, metrics in eval_results['per_class'].items():
print(f"{class_name}:")
print(f" Precision: {metrics['precision'] * 100:.2f}%")
print(f" Recall: {metrics['recall'] * 100:.2f}%")
print(f" F1 score: {metrics['f1'] * 100:.2f}%")
print(f" Samples: {metrics['support']}")
# Plot confusion matrix (if exists)
cm_path = Path('../results/confusion_matrix.png')
if cm_path.exists():
from IPython.display import Image
display(Image(str(cm_path)))
else:
print("Evaluation results file not found, please evaluate the model first")
# Load evaluation results (if exists) eval_path = Path('../results/evaluation_results.json') if eval_path.exists(): with open(eval_path, 'r', encoding='utf-8') as f: eval_results = json.load(f) print("=" * 70) print("Model evaluation results") print("=" * 70) print(f"Test set accuracy: {eval_results['accuracy'] * 100:.2f}%") print("\nMacro average:") print(f" Precision: {eval_results['precision_macro'] * 100:.2f}%") print(f" Recall: {eval_results['recall_macro'] * 100:.2f}%") print(f" F1 score: {eval_results['f1_macro'] * 100:.2f}%") print("\nDetailed metrics for each class:") print("-" * 70) for class_name, metrics in eval_results['per_class'].items(): print(f"{class_name}:") print(f" Precision: {metrics['precision'] * 100:.2f}%") print(f" Recall: {metrics['recall'] * 100:.2f}%") print(f" F1 score: {metrics['f1'] * 100:.2f}%") print(f" Samples: {metrics['support']}") # Plot confusion matrix (if exists) cm_path = Path('../results/confusion_matrix.png') if cm_path.exists(): from IPython.display import Image display(Image(str(cm_path))) else: print("Evaluation results file not found, please evaluate the model first")
====================================================================== Model evaluation results ====================================================================== Test set accuracy: 91.84% Macro average: Precision: 91.43% Recall: 89.47% F1 score: 90.04% Detailed metrics for each class: ---------------------------------------------------------------------- atarayo: Precision: 90.00% Recall: 69.23% F1 score: 78.26% Samples: 13 yoasobi: Precision: 89.47% Recall: 100.00% F1 score: 94.44% Samples: 17 yorushika: Precision: 89.47% Recall: 91.89% F1 score: 90.67% Samples: 37 zutomayo: Precision: 96.77% Recall: 96.77% F1 score: 96.77% Samples: 31