You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
gesture_rec/ml_gestures/train.py

136 lines
6.5 KiB
Python

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

import pandas as pd
import numpy as np
import joblib
import json
import argparse
from sklearn.linear_model import LogisticRegression
from sklearn.neural_network import MLPClassifier
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split
from sklearn.metrics import classification_report, confusion_matrix, accuracy_score
def balance_data(df, target_classes, random_state=42):
"""
Балансирует только указанные классы до минимального размера среди них.
Класс 'none' (и любые другие) остаются без изменений.
"""
# Разделяем на целевые и остальные
target_df = df[df['class'].isin(target_classes)]
other_df = df[~df['class'].isin(target_classes)]
# Определяем минимальный размер среди целевых классов
class_counts = target_df['class'].value_counts()
min_count = class_counts.min()
# Балансируем каждый целевой класс
balanced_parts = []
for cls in target_classes:
cls_df = target_df[target_df['class'] == cls]
if len(cls_df) > min_count:
cls_df = cls_df.sample(n=min_count, random_state=random_state)
balanced_parts.append(cls_df)
balanced_target = pd.concat(balanced_parts, ignore_index=True)
# Объединяем с остальными данными (none и др.)
balanced_df = pd.concat([balanced_target, other_df], ignore_index=True)
return balanced_df
def train(csv_path, model_path, model_type='mlp', test_size=0.2, random_state=42, balance=False, target_classes=None):
# Загрузка данных
df = pd.read_csv(csv_path)
print(f"Total samples: {len(df)}")
# Балансировка
if balance:
if target_classes is None:
# По умолчанию балансируем все классы кроме 'none' (если есть)
all_classes = df['class'].unique()
target_classes = [c for c in all_classes if c != 'none']
if not target_classes:
raise ValueError("No target classes found (none is the only class).")
print(f"Balancing target classes: {target_classes}")
df = balance_data(df, target_classes, random_state)
print(f"After balancing: {len(df)} samples")
print(df['class'].value_counts())
# Разделение на признаки и метки
X = df.iloc[:, 1:].values.astype(np.float32)
y_labels = df.iloc[:, 0].values
class_names = sorted(df['class'].unique())
label_to_id = {label: i for i, label in enumerate(class_names)}
y = np.array([label_to_id[label] for label in y_labels])
print(f"Classes: {class_names}")
print(f"Feature count: {X.shape[1]}")
# Стратифицированное разбиение
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=test_size, random_state=random_state, stratify=y
)
print(f"Train size: {len(X_train)}, Test size: {len(X_test)}")
# Выбор модели
if model_type == 'linear': #попробовать добавить логистическую регрессию с полиномиальными признаками, можно чисто для сравнения
model = LogisticRegression(max_iter=1000, random_state=random_state)
elif model_type == 'mlp':
model = MLPClassifier(hidden_layer_sizes=(64, 32), activation='relu',
solver='adam', max_iter=500, random_state=random_state,
early_stopping=True, validation_fraction=0.2)
elif model_type == 'rf':
model = RandomForestClassifier(n_estimators=50, max_depth=10, random_state=random_state)
else:
raise ValueError("model_type должен быть linear, mlp или rf")
# Обучение
model.fit(X_train, y_train)
# Предсказание на тесте
y_pred = model.predict(X_test)
# Метрики
accuracy = accuracy_score(y_test, y_pred)
report = classification_report(y_test, y_pred, target_names=class_names, output_dict=True)
conf_matrix = confusion_matrix(y_test, y_pred).tolist()
print(f"\nAccuracy: {accuracy:.4f}")
print("\nClassification Report:")
for cls in class_names:
print(f"{cls}: precision={report[cls]['precision']:.3f}, recall={report[cls]['recall']:.3f}, f1={report[cls]['f1-score']:.3f}")
print("\nConfusion Matrix:")
print(conf_matrix)
# Сохранение модели и отчёта
joblib.dump({'model': model, 'class_names': class_names}, model_path)
report_data = {
'model_type': model_type,
'accuracy': accuracy,
'classification_report': report,
'confusion_matrix': conf_matrix,
'train_samples': len(X_train),
'test_samples': len(X_test),
'classes': class_names,
'balance': balance,
'target_classes': target_classes if balance else None,
}
with open(model_path.replace('.pkl', '_report.json'), 'w') as f:
json.dump(report_data, f, indent=2)
print(f"\nModel saved to {model_path}")
print(f"Report saved to {model_path.replace('.pkl', '_report.json')}")
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Обучение модели для распознавания жестов')
parser.add_argument('--csv', required=True, help='Путь к CSV-файлу с данными')
parser.add_argument('--model', required=True, help='Путь для сохранения модели (.pkl)')
parser.add_argument('--type', default='mlp', choices=['linear', 'mlp', 'rf'], help='Тип модели')
parser.add_argument('--test_size', type=float, default=0.2, help='Доля тестовой выборки')
parser.add_argument('--random_state', type=int, default=42, help='Seed для воспроизводимости')
parser.add_argument('--balance', action='store_true', help='Балансировать классы (кроме none)')
parser.add_argument('--target_classes', type=str, default=None,
help='Список целевых классов для балансировки через запятую (по умолчанию все, кроме none)')
args = parser.parse_args()
target_classes = args.target_classes.split(',') if args.target_classes else None
train(args.csv, args.model, args.type, args.test_size, args.random_state, args.balance, target_classes)