init
This commit is contained in:
+203
@@ -0,0 +1,203 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# File Name : PolyBox.py
|
||||
# Created By : Yannis Bendi-Ouis
|
||||
|
||||
from StimulationsCodes import *
|
||||
from PolyStimulations import *
|
||||
from openvibe import *
|
||||
import sys, traceback, collections
|
||||
from io import StringIO
|
||||
import re
|
||||
|
||||
|
||||
def get_label_from_stim(stim):
|
||||
stims_1 = list(OpenViBE_stimulation.items())
|
||||
stims_2 = list(Poly_stimulation.items())
|
||||
inv_dictstim = {v: k for k, v in stims_1 + stims_2}
|
||||
label = inv_dictstim[stim.identifier][7:].lower()
|
||||
label = ' '.join(label.split('_'))
|
||||
return label
|
||||
|
||||
|
||||
class PolyBox(OVBox):
|
||||
|
||||
def __init__(self, record=True):
|
||||
OVBox.__init__(self)
|
||||
self.acquiring_channel = []
|
||||
self.signalHeader = []
|
||||
self.data = {}
|
||||
self.record = record
|
||||
|
||||
self.mode = ''
|
||||
self.current_stimulation = None
|
||||
self.labels = []
|
||||
|
||||
def initialize(self):
|
||||
|
||||
def verify_entry(self):
|
||||
# Verify that the inputs corresponds to ov-mode or poly-mode
|
||||
# ov-mode : 1 stimulations and 1 streamed matrix
|
||||
# poly-mode : several streamed-matrix
|
||||
# others : others -> warning
|
||||
list_entry_type = [entry.type() for entry in self.input]
|
||||
nb_matrix = list_entry_type.count('StreamedMatrix')
|
||||
nb_stim = list_entry_type.count('Stimulations')
|
||||
nb_signal = list_entry_type.count('Signal')
|
||||
|
||||
if nb_stim == 1 and nb_matrix == 1 and nb_signal == 0:
|
||||
# ov-mode
|
||||
self.mode = 'ov-mode'
|
||||
elif nb_stim == 0 and nb_matrix >= 1 and nb_signal == 0:
|
||||
# poly-mode
|
||||
self.mode = 'poly-mode'
|
||||
else:
|
||||
raise Exception("ERROR : Entry of the box does not corresponds to any mode. \
|
||||
You can use 1 stimulations and 1 streamed matrix, or several streamed matrix. \
|
||||
But you can not use {} StreamedMatrix, {} Stimulations and {} Signal entry.".format(nb_matrix, nb_stim, nb_signal))
|
||||
|
||||
def verify_stim_output(self):
|
||||
# Verify that an output stim exist, otherwise prevent the user that the program won't stop
|
||||
flag = False
|
||||
for out in self.output:
|
||||
if out.type() == 'Stimulations':
|
||||
flag = True
|
||||
break
|
||||
|
||||
if not flag:
|
||||
print('WARNING : The DatasetCreator does not have any output Stimulation. The program may never stop.')
|
||||
|
||||
def get_labels(self):
|
||||
# retrieve labels in form : label1, label2, label3, mon label4
|
||||
# Useless if you are in OV-MODE
|
||||
if 'Labels' in self.setting.keys():
|
||||
string = self.setting['Labels']
|
||||
if len(string) > 0:
|
||||
labels_cut = string.lower().split(',')
|
||||
for label in labels_cut:
|
||||
self.labels += ['_'.join([w for w in label.split(' ') if w != ''])]
|
||||
|
||||
def init_acquiring_channel(self):
|
||||
# We get data for every input channel
|
||||
for _ in range(len(self.input)):
|
||||
self.acquiring_channel += [False]
|
||||
self.signalHeader += [None]
|
||||
|
||||
verify_entry(self)
|
||||
verify_stim_output(self)
|
||||
get_labels(self)
|
||||
init_acquiring_channel(self)
|
||||
self.on_initialize()
|
||||
|
||||
def process(self):
|
||||
# we go through every input
|
||||
for inputIndex in range(len(self.input)):
|
||||
for chunkIndex in range(len(self.input[inputIndex])):
|
||||
|
||||
# Signal init
|
||||
if type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixHeader:
|
||||
self.header_received(inputIndex, chunkIndex)
|
||||
|
||||
# Process every chunk received
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixBuffer:
|
||||
self.chunk_received(inputIndex, chunkIndex)
|
||||
|
||||
# End of signal
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixEnd:
|
||||
self.end_received(inputIndex, chunkIndex)
|
||||
|
||||
# Stimulations init
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStimulationHeader:
|
||||
self.header_received(inputIndex, chunkIndex)
|
||||
|
||||
# Process every stimulation
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStimulationSet:
|
||||
self.stimulation_received(inputIndex, chunkIndex)
|
||||
|
||||
# End of stim
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStimulationEnd:
|
||||
self.end_received(inputIndex, chunkIndex)
|
||||
|
||||
def uninitialize(self):
|
||||
pass
|
||||
|
||||
# ------- * -------- * ---------
|
||||
|
||||
def header_received(self, inputIndex, chunkIndex):
|
||||
header = self.input[inputIndex].pop()
|
||||
self.signalHeader[inputIndex] = header
|
||||
self.acquiring_channel[inputIndex] = True
|
||||
self.on_header_received(header)
|
||||
|
||||
def chunk_received(self, inputIndex, chunkIndex):
|
||||
chunk = list(self.input[inputIndex].pop())
|
||||
if self.acquiring_channel[inputIndex]:
|
||||
|
||||
# We look for the best key to use in function of mode and settings labels.
|
||||
key = None
|
||||
if self.mode == 'poly-mode':
|
||||
if len(self.labels) > 0:
|
||||
key = self.labels[inputIndex]
|
||||
else:
|
||||
key = inputIndex
|
||||
|
||||
elif self.mode == 'ov-mode':
|
||||
key = get_label_from_stim(self.current_stimulation)
|
||||
|
||||
if self.record:
|
||||
try:
|
||||
self.data[key].append(chunk)
|
||||
except KeyError:
|
||||
self.data[key] = [chunk]
|
||||
|
||||
shape = tuple(self.signalHeader[inputIndex].dimensionSizes)
|
||||
self.on_chunk_received(chunk, key, shape)
|
||||
|
||||
def stimulation_received(self, inputIndex, chunkIndex):
|
||||
stim_list = self.input[inputIndex].pop()
|
||||
if len(stim_list) > 0:
|
||||
self.current_stimulation = stim_list[0]
|
||||
|
||||
def end_received(self, inputIndex, chunkIndex):
|
||||
self.acquiring_channel[inputIndex] = False
|
||||
self.input[inputIndex].pop()
|
||||
if not self.is_acquiring():
|
||||
print("Fin de l'acquisition des données...")
|
||||
self.on_end_box()
|
||||
self.send_end_stim()
|
||||
|
||||
def is_acquiring(self):
|
||||
# Return false when all inputs of type StreamedMatrix received End flag.
|
||||
nb_inputs = len(self.input)
|
||||
for i in range(nb_inputs):
|
||||
acquiring = self.acquiring_channel[i]
|
||||
kind = self.input[i].type()
|
||||
if kind == 'StreamedMatrix' and acquiring:
|
||||
return True
|
||||
return False
|
||||
|
||||
def send_end_stim(self):
|
||||
indice = -1
|
||||
for i, out in enumerate(self.output):
|
||||
if out.type() == 'Stimulations':
|
||||
indice = i
|
||||
|
||||
if indice != -1:
|
||||
stimLabel = 'OVTK_StimulationId_ExperimentStop'
|
||||
stimCode = OpenViBE_stimulation[stimLabel]
|
||||
stimSet = OVStimulationSet(0, self.getCurrentTime())
|
||||
stimSet.append(OVStimulation(stimCode, self.getCurrentTime(), 0.))
|
||||
self.output[indice].append(stimSet)
|
||||
|
||||
# ---------- * -------------- * ---------------
|
||||
|
||||
def on_initialize(self):
|
||||
pass
|
||||
|
||||
def on_header_received(self, header):
|
||||
pass
|
||||
|
||||
def on_chunk_received(self, chunk, label, shape):
|
||||
pass
|
||||
|
||||
def on_end_box(self):
|
||||
pass
|
||||
+18
@@ -0,0 +1,18 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
#File Name : StimulationsCodes.py
|
||||
#Created By :
|
||||
|
||||
## Stimulation codes
|
||||
# Avoid the declaration of new stimulation (added in new tab here). Only Stimulation contains gdf in original list is really standard...
|
||||
Poly_stimulation = {
|
||||
'OVPoly_Down' : 0x10001, #Useless use 'OVTK_GDF_Down'
|
||||
'OVPoly_Up' : 0x10002, #Useless use 'OVTK_GDF_Up'
|
||||
'OVPoly_Right' : 0x10003, #Useless use 'OVTK_GDF_Right'
|
||||
'OVPoly_Left' : 0x10004, #Useless use 'OVTK_GDF_Left'
|
||||
'OVPoly_Neutral' : 0x10005,
|
||||
'OVPoly_Push' : 0x10006,
|
||||
'OVPoly_Pull' : 0x10007,
|
||||
'OVPoly_Left_Wink' : 0x10008,
|
||||
'OVPoly_Right_Wink' : 0x10009,
|
||||
# <Flag> New Stims
|
||||
}
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from matplotlib import pyplot as plt
|
||||
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA
|
||||
from sklearn.decomposition import PCA
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pickle
|
||||
from PolyBox import PolyBox
|
||||
import warnings
|
||||
warnings.filterwarnings("ignore")
|
||||
|
||||
|
||||
flag_3D = True
|
||||
try:
|
||||
from mpl_toolkits.mplot3d import Axes3D
|
||||
except:
|
||||
print('Unable to import Axes3D from mpl_toolkits.mplot3d, 3D visualization disabled.')
|
||||
flag_3D = False
|
||||
|
||||
|
||||
class DataViz(PolyBox):
|
||||
|
||||
def __init__(self):
|
||||
PolyBox.__init__(self)
|
||||
|
||||
self.x_data = []
|
||||
self.y_data = []
|
||||
self.model = None
|
||||
|
||||
self.path_load_model = ''
|
||||
self.path_save_model = ''
|
||||
self.algo = ''
|
||||
self.dimension_reduction = -1
|
||||
|
||||
def on_initialize(self):
|
||||
# We retrieve the setting from OpenViBE
|
||||
|
||||
def retrieve_path_save_model(self):
|
||||
try:
|
||||
self.path_save_model = self.setting["Path to save the model"]
|
||||
except KeyError:
|
||||
pass
|
||||
if self.path_save_model == '':
|
||||
print('No path has been given to save the model, thus it won\'t be saved.')
|
||||
|
||||
def retrieve_path_load_model(self):
|
||||
try:
|
||||
self.path_load_model = self.setting["Path to load the model"]
|
||||
except KeyError:
|
||||
pass
|
||||
if self.path_load_model == '':
|
||||
print('No path has been given to load the model, thus a new model will be created.')
|
||||
|
||||
def retrieve_algo(self):
|
||||
try:
|
||||
self.algo = self.setting['Algorithm (PCA or LDA)'].upper()
|
||||
except KeyError:
|
||||
pass
|
||||
if self.algo == '':
|
||||
print('No algo has been given to the model, default algo is PCA.')
|
||||
self.algo = 'PCA'
|
||||
|
||||
def retrieve_dimension_reduction(self):
|
||||
try:
|
||||
self.dimension_reduction = int(self.setting['Dimension reduction'])
|
||||
except KeyError:
|
||||
pass
|
||||
except ValueError:
|
||||
print('{} is not a number. Please use 2 or 3 dimensions only.'.format(self.setting['Dimension reduction']))
|
||||
|
||||
if self.dimension_reduction == -1:
|
||||
print('No dimension reduction has been given, default value is 2.')
|
||||
self.dimension_reduction = 2
|
||||
|
||||
elif self.dimension_reduction == 3 and not flag_3D:
|
||||
print('3D disabled, cannot show data in 3 dimensions. Default dimension is 2.')
|
||||
self.dimension_reduction = 3
|
||||
|
||||
retrieve_path_save_model(self)
|
||||
retrieve_path_load_model(self)
|
||||
retrieve_algo(self)
|
||||
retrieve_dimension_reduction(self)
|
||||
def on_end_box(self):
|
||||
self.prepare_data()
|
||||
self.make_model_and_transform()
|
||||
self.make_plot()
|
||||
|
||||
# ---------
|
||||
|
||||
def prepare_data(self):
|
||||
for label in self.data.keys():
|
||||
self.x_data += self.data[label]
|
||||
self.y_data += [label for _ in range(len(self.data[label]))]
|
||||
|
||||
def make_model_and_transform(self):
|
||||
# load the model if it exists, else create a new one
|
||||
# then transform the data
|
||||
|
||||
def load_model(self):
|
||||
model = pickle.load(open(self.path_load_model, 'rb'))
|
||||
print('Model load from {}.'.format(self.path_load_model))
|
||||
return model
|
||||
|
||||
def save_model(self):
|
||||
pickle.dump(self.model, open(self.path_save_model, 'wb'))
|
||||
print('Dataviz model saved in {}'.format(self.path_save_model))
|
||||
|
||||
def map_algo(self):
|
||||
switcher = { 'LDA': LDA, 'PCA': PCA }
|
||||
clf = switcher.get(self.algo)
|
||||
return clf, switcher
|
||||
|
||||
def create_fit_model(self):
|
||||
clf, _ = map_algo(self)
|
||||
clf = clf(n_components=self.dimension_reduction)
|
||||
if self.algo == 'PCA':
|
||||
clf.fit(self.x_data)
|
||||
elif self.algo == 'LDA':
|
||||
clf.fit(self.x_data, self.y_data)
|
||||
else:
|
||||
raise Exception('{} is not known as an Algorithm. Please use PCA or LDA.'.format(self.algo))
|
||||
return clf
|
||||
|
||||
# Load or create the model
|
||||
if len(self.path_load_model) > 0:
|
||||
self.model = load_model(self)
|
||||
else:
|
||||
self.model = create_fit_model(self)
|
||||
|
||||
# Save the model
|
||||
if self.path_save_model != '':
|
||||
save_model(self)
|
||||
|
||||
# Transform data
|
||||
self.x_data = self.model.transform(self.x_data)
|
||||
self.y_data = np.array(self.y_data)
|
||||
|
||||
def make_plot(self):
|
||||
|
||||
fig = plt.figure(figsize=(12, 12))
|
||||
|
||||
all_labels = list(self.data.keys())
|
||||
colors = np.array([all_labels.index(label) for label in self.y_data])
|
||||
|
||||
if self.dimension_reduction == 2:
|
||||
|
||||
ax = plt.axes()
|
||||
|
||||
for label in all_labels:
|
||||
ax.text(self.x_data[self.y_data == label, 0].mean(), self.x_data[self.y_data == label, 1].mean(),
|
||||
label, horizontalalignment='center', bbox=dict(alpha=0.5, edgecolor='w', facecolor='w'))
|
||||
|
||||
ax.scatter(self.x_data[:, 0], self.x_data[:, 1], alpha=0.5, c=colors, cmap='Spectral', edgecolor='g')
|
||||
|
||||
plt.show()
|
||||
elif self.dimension_reduction == 3:
|
||||
|
||||
ax = Axes3D(fig)
|
||||
|
||||
for label in all_labels:
|
||||
ax.text3D(self.x_data[self.y_data == label, 0].mean(),
|
||||
self.x_data[self.y_data == label, 1].mean(),
|
||||
self.x_data[self.y_data == label, 2].mean(),
|
||||
label, horizontalalignment='center', bbox=dict(alpha=0.5, edgecolor='w', facecolor='w'))
|
||||
|
||||
ax.scatter(self.x_data[:, 0], self.x_data[:, 1], self.x_data[:, 2],
|
||||
alpha=0.5, c=colors, cmap='Spectral', edgecolor='g')
|
||||
|
||||
plt.show()
|
||||
|
||||
|
||||
box = DataViz()
|
||||
|
||||
+421
@@ -0,0 +1,421 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
'''
|
||||
* Software License Agreement (AGPL-3 License)
|
||||
*
|
||||
* OpenViBE
|
||||
* Copyright (C) Inria, 2006-2019
|
||||
*
|
||||
* Authors
|
||||
*
|
||||
* 2019, Yannis Bendi-Ouis <yannis.bendiouis@gmail.com>
|
||||
* 2019, Jimmy Leblanc <jimmy.leblanc01@gmail.com>
|
||||
*
|
||||
* This program is free software: you can redistribute it and/or modify
|
||||
* it under the terms of the GNU Affero General Public License version 3,
|
||||
* as published by the Free Software Foundation.
|
||||
*
|
||||
* This program is distributed in the hope that it will be useful,
|
||||
* but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
* GNU Affero General Public License for more details.
|
||||
*
|
||||
* You should have received a copy of the GNU Affero General Public License
|
||||
* along with this program.
|
||||
* If not, see <http://www.gnu.org/licenses/>.
|
||||
*
|
||||
'''
|
||||
|
||||
from pandas import Series, DataFrame, read_csv
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from natsort import natsorted
|
||||
from PolyStimulations import Poly_stimulation
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import inspect
|
||||
|
||||
# Dans les settings :
|
||||
# - "Path directory" : path qui mène au directory de sauvegarde des données. Donc finir par un '/'.
|
||||
# - "Label_X" : Nom des labels, où X est un nombre. Il en faut autant qu'il y a de labels.
|
||||
# - "Several CSV" : Boolean ou string qui return "true" ou "false".
|
||||
# - "Number of folds" : Integer qui indique en combien de différents folds les data doivent être séparées.
|
||||
# - "Number of actions" : Integer qui indique le nombre d'actions à enregistrer lors d'une session.
|
||||
|
||||
#------------------------------------------------------------
|
||||
ACTION_DURATION = 12 # in seconds
|
||||
TAMPON_DURATION = 3 # in seconds
|
||||
BEGIN_RECORD = 2 # in seconds
|
||||
END_RECORD = ACTION_DURATION
|
||||
FREQ = 128
|
||||
NB_POINTS_ONE_RECORD = ACTION_DURATION * FREQ
|
||||
NB_POINTS_ONE_TAMPON = TAMPON_DURATION * FREQ
|
||||
|
||||
CHANNELS_NAME = ['AF3', 'F7', 'F3', 'FC5', 'T7', 'P7', 'O1', 'O2', 'P8', 'T8', 'FC6', 'F4', 'F8', 'AF4'] # 14
|
||||
CHANNELS_STIMULATION = ['Event Id', 'Event Date', 'Event Duration']
|
||||
|
||||
PREFIXE_STIM = 'OVPoly_'
|
||||
|
||||
#------------------------------------------------------------
|
||||
def ovdf(df, rewrite_stim=True):
|
||||
shape = df.shape
|
||||
timer = Series([float(i)/FREQ for i in range(shape[0])])
|
||||
epochs = Series([int(i/64) for i in range(shape[0])])
|
||||
|
||||
df.insert(0, 'Time:{}Hz'.format(FREQ), timer)
|
||||
df.insert(1, 'Epoch', epochs)
|
||||
|
||||
if rewrite_stim:
|
||||
df['Event Id'] = Series([None for i in range(shape[0])])
|
||||
df['Event Date'] = Series([None for i in range(shape[0])])
|
||||
df['Event Duration'] = Series([None for i in range(shape[0])])
|
||||
|
||||
else:
|
||||
# re-compute the time informations for stimulations
|
||||
stim = df['Event Id'].values.tolist()
|
||||
date = []
|
||||
duration = []
|
||||
for i, event_id in enumerate(stim):
|
||||
if not pd.isnull(event_id):
|
||||
date += [timer[i]]
|
||||
duration += [0]
|
||||
else:
|
||||
date += [None]
|
||||
duration += [None]
|
||||
df['Event Date'] = Series(date)
|
||||
df['Event Duration'] = Series(duration)
|
||||
|
||||
|
||||
#------------------------------------------------------------
|
||||
def get_stim_code_from_label(label):
|
||||
key = label.split('_')
|
||||
key = "_".join([w[0].upper() + w[1:] for w in key])
|
||||
key = PREFIXE_STIM + key
|
||||
return Poly_stimulation[key]
|
||||
|
||||
|
||||
#------------------------------------------------------------
|
||||
class DatasetCreator(OVBox):
|
||||
|
||||
#----------------------------------------
|
||||
def __init__(self):
|
||||
OVBox.__init__(self)
|
||||
self.dir_name = None
|
||||
self.url = None
|
||||
self.signalHeader = None
|
||||
self.labels_queue = None
|
||||
self.current_label = None
|
||||
|
||||
self.is_tampon = False
|
||||
self.several_csv = False
|
||||
|
||||
self.nb_fold = 0
|
||||
self.nb_action_per_session = 0
|
||||
|
||||
self.labels = []
|
||||
self.data_recorded = []
|
||||
self.labels_order = []
|
||||
self.dic_fold = {}
|
||||
self.dic_dicount = {}
|
||||
|
||||
#----------------------------------------
|
||||
def initialize(self):
|
||||
|
||||
def verify_labels_correct(self):
|
||||
for label in self.labels:
|
||||
try:
|
||||
get_stim_code_from_label(label)
|
||||
except KeyError:
|
||||
raise Exception('Label {} not defined in PolyStimulations. You may want to add it with the manager.'.format(label))
|
||||
|
||||
def get_labels(self):
|
||||
param_names = self.setting.keys()
|
||||
param_names = natsorted(param_names)
|
||||
for n in param_names:
|
||||
if 'Label_' in n:
|
||||
label = self.setting[n]
|
||||
label = label.replace(' ', '_')
|
||||
if len(label) > 0:
|
||||
self.labels += [label.lower()]
|
||||
|
||||
def init_dict(self):
|
||||
self.dic_fold = {'fold_{}'.format(i): None for i in range(1, self.nb_fold+1)}
|
||||
self.dic_dicount = {'fold_{}'.format(i): None for i in range(1, self.nb_fold+1)}
|
||||
|
||||
def retrieve_settings(self):
|
||||
|
||||
# On récupère la booleen indiquant si l'on souhaite plusieurs ou un seul csv par fold
|
||||
self.several_csv = self.setting['Several CSV']
|
||||
if self.several_csv == 'true':
|
||||
self.several_csv = True
|
||||
elif self.several_csv == 'false':
|
||||
self.several_csv = False
|
||||
# On récupère le path du directory où l'on créé les fold
|
||||
self.dir_name = self.setting['Path directory']
|
||||
if self.dir_name[-1] != '/':
|
||||
self.dir_name += '/'
|
||||
|
||||
# On récupère le nombre de fold
|
||||
self.nb_fold = int(self.setting['Number of folds'])
|
||||
# On récupère le nombre d'action à record par session
|
||||
self.nb_action_per_session = int(self.setting['Number of actions'])
|
||||
# On récupère les labels
|
||||
get_labels(self)
|
||||
|
||||
def verify_stim_output(self):
|
||||
# Verify that an output stim exist, otherwise prevent the user that the program won't stop
|
||||
flag = False
|
||||
for out in self.output:
|
||||
if out.type() == 'Stimulations':
|
||||
flag = True
|
||||
break
|
||||
|
||||
if not flag:
|
||||
print('WARNING : The DatasetCreator does not have any output Stimulation. The program may never stop.')
|
||||
|
||||
retrieve_settings(self)
|
||||
verify_stim_output(self)
|
||||
verify_labels_correct(self)
|
||||
init_dict(self)
|
||||
|
||||
self.verify_arborescence()
|
||||
self.prepare_session()
|
||||
self.new_record()
|
||||
|
||||
#----------------------------------------
|
||||
def process(self):
|
||||
# On parcours tous les inputs
|
||||
for inputIndex in range(len(self.input)):
|
||||
for chunkIndex in range(len(self.input[inputIndex])):
|
||||
|
||||
# Initialisation pour le signal
|
||||
if type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixHeader:
|
||||
self.header_received(inputIndex, chunkIndex)
|
||||
|
||||
# Traitement à effectuer pour chaque chunk reçu
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixBuffer:
|
||||
self.chunk_received(inputIndex, chunkIndex)
|
||||
|
||||
# Fin du signal
|
||||
elif type(self.input[inputIndex][chunkIndex]) == OVStreamedMatrixEnd:
|
||||
self.end_received(inputIndex, chunkIndex)
|
||||
|
||||
#----------------------------------------
|
||||
def uninitialize(self):
|
||||
pass
|
||||
|
||||
# -------------- * ------------- * -------------
|
||||
|
||||
#----------------------------------------
|
||||
def header_received(self, inputIndex, chunkIndex):
|
||||
self.signalHeader = self.input[inputIndex].pop()
|
||||
|
||||
#----------------------------------------
|
||||
def chunk_received(self, inputIndex, chunkIndex):
|
||||
chunk = self.input[inputIndex].pop()
|
||||
|
||||
# Plusieurs lignes sont envoyées en même temps, il faut les séparer
|
||||
indices = [(len(CHANNELS_NAME)*i, len(CHANNELS_NAME)*(i+1)) for i in range(int(len(chunk)/len(CHANNELS_NAME)))]
|
||||
for begin, end in indices:
|
||||
self.data_recorded += [chunk[begin:end]]
|
||||
|
||||
# Fin d'un tampon
|
||||
if self.is_tampon and len(self.data_recorded) >= NB_POINTS_ONE_TAMPON:
|
||||
self.empty_tampon()
|
||||
self.new_record()
|
||||
|
||||
# Fin d'un record
|
||||
if not self.is_tampon and len(self.data_recorded) >= NB_POINTS_ONE_RECORD:
|
||||
if not self.end_record():
|
||||
self.end_of_session()
|
||||
self.end_of_box()
|
||||
|
||||
#----------------------------------------
|
||||
def end_received(self, inputIndex, chunkIndex):
|
||||
self.input[inputIndex].pop()
|
||||
print("Error : entry signal stoped.")
|
||||
self.end_of_box()
|
||||
|
||||
#----------------------------------------
|
||||
def end_of_box(self):
|
||||
indice = -1
|
||||
for i, out in enumerate(self.output):
|
||||
if out.type() == 'Stimulations':
|
||||
indice = i
|
||||
|
||||
if indice != -1:
|
||||
stimLabel = 'OVTK_StimulationId_ExperimentStop'
|
||||
stimCode = OpenViBE_stimulation[stimLabel]
|
||||
stimSet = OVStimulationSet(0, self.getCurrentTime())
|
||||
stimSet.append(OVStimulation(stimCode, self.getCurrentTime(), 0.))
|
||||
self.output[0].append(stimSet)
|
||||
|
||||
# -------------- * ------------- * -------------
|
||||
|
||||
#----------------------------------------
|
||||
def verify_arborescence(self):
|
||||
# On vérifie que chaque dossier du path existe, sinon on les créé
|
||||
# et on vérifie les dicount
|
||||
|
||||
def verify_dir(self):
|
||||
# On vérifie que le dossier existe et ses sous-dossiers, sinon on les créé
|
||||
if not os.path.exists(self.dir_name):
|
||||
os.mkdir(self.dir_name, 0o775)
|
||||
|
||||
def verify_folds(self):
|
||||
# On vérifie que les dossiers des différents folds existent, sinon on les créés
|
||||
for i in range(1, self.nb_fold+1):
|
||||
name = self.dir_name + 'fold_{}/'.format(i)
|
||||
if not os.path.exists(name):
|
||||
os.mkdir(name, 0o775)
|
||||
|
||||
def verify_and_load_dicount(self):
|
||||
# On vérifie que les fichiers contenant les compteurs par label existent, sinon on les créé
|
||||
for i in range(1, self.nb_fold+1):
|
||||
fold = 'fold_{}'.format(i)
|
||||
filename = self.dir_name + fold + '/dicount.pick'
|
||||
|
||||
if not os.path.isfile(filename):
|
||||
dicount = {l: 0 for l in self.labels}
|
||||
self.dic_dicount[fold] = dicount
|
||||
pickle.dump(dicount, open(filename, 'wb'))
|
||||
else:
|
||||
self.dic_dicount[fold] = pickle.load(open(filename, 'rb'))
|
||||
|
||||
verify_dir(self)
|
||||
verify_folds(self)
|
||||
verify_and_load_dicount(self)
|
||||
|
||||
#----------------------------------------
|
||||
def prepare_session(self):
|
||||
|
||||
def create_labels_queue(self):
|
||||
nb_label = len(self.labels)
|
||||
nb_total = int(self.nb_action_per_session / nb_label) # in Python 3 / create float automatically and range hate that
|
||||
nb_reste = self.nb_action_per_session - nb_total*nb_label
|
||||
|
||||
choice_label = []
|
||||
tmp = [l for l in self.labels]
|
||||
for _ in range(nb_reste):
|
||||
l = random.choice(tmp)
|
||||
tmp.remove(l)
|
||||
choice_label += [l]
|
||||
|
||||
tmp = [l for l in self.labels]
|
||||
self.labels_queue = tmp*nb_total + choice_label
|
||||
random.shuffle(self.labels_queue)
|
||||
|
||||
# Prepare la suite aléatoire d'action a exectuer
|
||||
create_labels_queue(self)
|
||||
|
||||
# Initialise self.dic_fold
|
||||
for key in self.dic_fold.keys():
|
||||
self.dic_fold[key] = {l: [] for l in self.labels}
|
||||
|
||||
#----------------------------------------
|
||||
def empty_tampon(self):
|
||||
# Supprime les données tampons entre deux records
|
||||
self.is_tampon = False
|
||||
end = TAMPON_DURATION * FREQ
|
||||
self.data_recorded = self.data_recorded[end:]
|
||||
|
||||
#----------------------------------------
|
||||
def new_record(self):
|
||||
# Prevent the user a new record begin
|
||||
if len(self.labels_queue) > 0:
|
||||
self.current_label = self.labels_queue.pop(0)
|
||||
print('Current label : {}'.format(self.current_label))
|
||||
|
||||
#----------------------------------------
|
||||
def end_record(self):
|
||||
# Démarre un nouvel enregistrement de 10 sec. Si le dernier est fini, l'enregistre.
|
||||
# Return True s'il y a encore d'autres label à étudié pour la session, False sinon.
|
||||
|
||||
def retrieve_record(self):
|
||||
|
||||
def add_to_fold(self, data, label):
|
||||
# Ajoute les éléments dans la liste data au dic_fold label et mets à jour dic_dicount
|
||||
print('hello', self.dic_dicount)
|
||||
tmp = [(key, value[label]) for key, value in self.dic_dicount.items()]
|
||||
fold = min(tmp, key=lambda x: x[1])[0]
|
||||
self.dic_fold[fold][label] += [data]
|
||||
self.dic_dicount[fold][label] += len(data)
|
||||
self.labels_order += [label]
|
||||
|
||||
# Extract the 10s recorded
|
||||
record = self.data_recorded[:NB_POINTS_ONE_RECORD]
|
||||
self.data_recorded = self.data_recorded[NB_POINTS_ONE_RECORD:]
|
||||
|
||||
# Extract data between begin and end
|
||||
begin = FREQ*BEGIN_RECORD
|
||||
end = FREQ*END_RECORD
|
||||
data = record[begin:end]
|
||||
add_to_fold(self, data, self.current_label)
|
||||
|
||||
# End record
|
||||
print('Stop.')
|
||||
|
||||
retrieve_record(self)
|
||||
self.is_tampon = True
|
||||
|
||||
return len(self.labels_queue) > 0
|
||||
|
||||
#----------------------------------------
|
||||
def end_of_session(self):
|
||||
|
||||
def maj_dicount(self, fold):
|
||||
filename = self.dir_name + fold + '/dicount.pick'
|
||||
dicount = self.dic_dicount[fold]
|
||||
pickle.dump(dicount, open(filename, 'wb'))
|
||||
|
||||
def add_label_stimulation(self, data, label):
|
||||
# Add the stimulations indicating the begginning of a label
|
||||
length = len(data)
|
||||
event_id = [None for _ in range(length)]
|
||||
event_date = [None for _ in range(length)]
|
||||
event_duration = [None for _ in range(length)]
|
||||
event_id[0] = get_stim_code_from_label(label)
|
||||
|
||||
dict_event = {"Event Id": event_id, "Event Date": event_date, "Event Duration": event_duration}
|
||||
dict_data = {c: np.array(data)[:, i].tolist() for i, c in enumerate(CHANNELS_NAME)}
|
||||
dict_data.update(dict_event)
|
||||
|
||||
return dict_data
|
||||
|
||||
def append_data_to_csv(self, data, fold, label, several_csv=False):
|
||||
# Create all csv, either you can us one CSV with stimulations, either one CSV per label
|
||||
if several_csv:
|
||||
filename = self.dir_name + fold + '/' + label + '.csv'
|
||||
columns = CHANNELS_NAME
|
||||
else:
|
||||
filename = self.dir_name + fold + '/' + fold + '.csv'
|
||||
columns = CHANNELS_NAME + CHANNELS_STIMULATION
|
||||
data = add_label_stimulation(self, data, label)
|
||||
|
||||
old_df = DataFrame()
|
||||
if os.path.isfile(filename):
|
||||
old_df = read_csv(filename).filter(columns)
|
||||
|
||||
add_df = DataFrame(data, columns=columns)
|
||||
new_df = old_df.append(add_df, ignore_index=True)
|
||||
|
||||
ovdf(new_df, rewrite_stim=several_csv)
|
||||
new_df.to_csv(filename, index=False)
|
||||
labels_count = {label: 0 for label in self.labels}
|
||||
|
||||
self.labels_order = {'fold_{}'.format(n): self.labels_order[self.nb_action_per_session*i: self.nb_action_per_session*(i+1)]
|
||||
for i in range(self.nb_fold) for n in range(1, self.nb_fold+1)}
|
||||
|
||||
for i in range(1, self.nb_fold+1):
|
||||
fold = 'fold_{}'.format(i)
|
||||
|
||||
for label in self.labels_order[fold]:
|
||||
count = labels_count[label]
|
||||
data = self.dic_fold[fold][label][count]
|
||||
labels_count[label] += 1
|
||||
append_data_to_csv(self, data, fold, label, several_csv=self.several_csv)
|
||||
|
||||
maj_dicount(self, fold)
|
||||
|
||||
|
||||
box = DatasetCreator()
|
||||
+107
@@ -0,0 +1,107 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from pyriemann.classification import MDM
|
||||
from pyriemann.tangentspace import TangentSpace
|
||||
from pyriemann.estimation import Covariances
|
||||
from sklearn.pipeline import make_pipeline
|
||||
from sklearn.tree import DecisionTreeClassifier
|
||||
from sklearn.linear_model import LogisticRegression, SGDClassifier
|
||||
from sklearn.neural_network import MLPClassifier
|
||||
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis
|
||||
from sklearn.naive_bayes import GaussianNB
|
||||
from sklearn.svm import SVC
|
||||
from sklearn.ensemble import RandomForestClassifier, AdaBoostClassifier, ExtraTreesClassifier, BaggingClassifier
|
||||
from sklearn.neighbors import KNeighborsClassifier, NearestCentroid
|
||||
from sklearn.metrics import confusion_matrix, classification_report
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
import pickle
|
||||
import os
|
||||
from PolyBox import PolyBox
|
||||
import warnings
|
||||
warnings.filterwarnings("ignore")
|
||||
|
||||
|
||||
class ProcessML(PolyBox):
|
||||
|
||||
#----------------------------------------
|
||||
def __init__(self):
|
||||
PolyBox.__init__(self, record=False)
|
||||
self.model_path = None
|
||||
self.pred_saving_path = None
|
||||
self.model = None
|
||||
self.predictions = []
|
||||
self.shape = None
|
||||
|
||||
#----------------------------------------
|
||||
def on_initialize(self):
|
||||
# we get the model file
|
||||
self.model_path = self.setting['Model filename']
|
||||
|
||||
# we load the model
|
||||
try:
|
||||
self.model = pickle.load(open(self.model_path, 'rb'))
|
||||
except IOError as err:
|
||||
print(err)
|
||||
print('Please indicate an existing model.')
|
||||
self.send_end_stim()
|
||||
|
||||
try:
|
||||
self.pred_saving_path = self.setting['Predictions filename']
|
||||
if len(self.pred_saving_path.replace(' ', '')) > 0:
|
||||
self.pred_saving_path = os.path.abspath(self.pred_saving_path)
|
||||
except IOError:
|
||||
print('No filename to save predictions has been given, they will not be saved')
|
||||
|
||||
#----------------------------------------
|
||||
def on_end_box(self):
|
||||
self.save_preds()
|
||||
self.make_stats()
|
||||
|
||||
#----------------------------------------
|
||||
def list_to_str(self, preds):
|
||||
string = ''
|
||||
for x in preds:
|
||||
string += str(x) + ','
|
||||
return string
|
||||
|
||||
#----------------------------------------
|
||||
def save_preds(self):
|
||||
|
||||
if self.pred_saving_path is not None and self.pred_saving_path != "":
|
||||
|
||||
preds_str = self.list_to_str(self.predictions)
|
||||
|
||||
with open(self.pred_saving_path, 'wb') as file:
|
||||
file.write(preds_str)
|
||||
print('Predictions saved in {}\n'.format(self.pred_saving_path))
|
||||
|
||||
#----------------------------------------
|
||||
def on_chunk_received(self, chunk, label, shape):
|
||||
|
||||
chunk = np.array(chunk)
|
||||
|
||||
if self.shape is None:
|
||||
# Riemanian Geometry
|
||||
if self.model.custom_classifier == 'Riemann Minimum Distance to Mean' or self.model.custom_classifier == 'Riemann Tangent Space':
|
||||
self.shape = (1, shape[0], shape[1])
|
||||
else:
|
||||
self.shape = (1, chunk.shape[0])
|
||||
|
||||
chunk = chunk.reshape(self.shape)
|
||||
pred = self.model.predict(chunk)
|
||||
|
||||
self.predictions += list(pred)
|
||||
|
||||
#----------------------------------------
|
||||
def make_stats(self):
|
||||
|
||||
length = len(self.predictions)
|
||||
dictpred = defaultdict(int)
|
||||
|
||||
for elem in self.predictions:
|
||||
dictpred[elem] += 1
|
||||
print("Metrics : \n")
|
||||
print('\n'.join(['{} : {}'.format(l, float(v)/length) for l, v in dictpred.items()]))
|
||||
|
||||
|
||||
box = ProcessML()
|
||||
+247
@@ -0,0 +1,247 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from pyriemann.classification import MDM
|
||||
from pyriemann.tangentspace import TangentSpace
|
||||
from pyriemann.estimation import Covariances
|
||||
import time
|
||||
from sklearn.pipeline import make_pipeline
|
||||
from sklearn.tree import DecisionTreeClassifier
|
||||
from sklearn.linear_model import LogisticRegression, SGDClassifier
|
||||
from sklearn.neural_network import MLPClassifier
|
||||
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis
|
||||
from sklearn.naive_bayes import GaussianNB
|
||||
from sklearn.svm import LinearSVC
|
||||
from sklearn.ensemble import RandomForestClassifier, AdaBoostClassifier, ExtraTreesClassifier, BaggingClassifier
|
||||
from sklearn.neighbors import KNeighborsClassifier, NearestCentroid
|
||||
from sklearn.metrics import confusion_matrix, classification_report
|
||||
from natsort import natsorted
|
||||
import random
|
||||
import pickle
|
||||
import numpy as np
|
||||
from collections import defaultdict
|
||||
from PolyBox import PolyBox
|
||||
import warnings
|
||||
warnings.filterwarnings("ignore")
|
||||
|
||||
|
||||
#----------------------------------------
|
||||
def train_test_split(data_dict, test_size=0.2):
|
||||
def create_xy(data_dict):
|
||||
# used to shuffle data in a way that prevent any training data to be in the test set
|
||||
y = np.array([])
|
||||
for label in data_dict:
|
||||
y_tmp = np.array([label for _ in data_dict[label]])
|
||||
y = np.concatenate((y, y_tmp))
|
||||
|
||||
array = []
|
||||
for label in data_dict:
|
||||
array.append(data_dict[label])
|
||||
|
||||
x = np.concatenate(array, axis=0)
|
||||
# shuffling
|
||||
permutation = np.random.permutation(x.shape[0])
|
||||
|
||||
permut = True
|
||||
if permut:
|
||||
x = x[permutation]
|
||||
y = y[permutation]
|
||||
|
||||
return x, y
|
||||
|
||||
dict_train = {}
|
||||
dict_test = {}
|
||||
|
||||
for key, data in data_dict.items():
|
||||
icut = int(len(data)*test_size)
|
||||
dict_test[key] = data[:icut]
|
||||
dict_train[key] = data[icut:]
|
||||
|
||||
x_train, y_train = create_xy(dict_train)
|
||||
x_test, y_test = create_xy(dict_test)
|
||||
|
||||
return np.array(x_train), np.array(y_train), np.array(x_test), np.array(y_test)
|
||||
|
||||
|
||||
#----------------------------------------
|
||||
class TrainerML(PolyBox):
|
||||
|
||||
#----------------------------------------
|
||||
# is exec first
|
||||
def __init__(self):
|
||||
PolyBox.__init__(self)
|
||||
self.model_path = None
|
||||
self.config_path = None
|
||||
self.save_path = None
|
||||
self.model = None
|
||||
self.clf = None
|
||||
self.config = None
|
||||
self.std_settings = []
|
||||
self.clf_dependant_settings = None
|
||||
|
||||
#----------------------------------------
|
||||
def on_initialize(self):
|
||||
|
||||
try:
|
||||
self.model_path = self.setting['Filename to save model to']
|
||||
except KeyError:
|
||||
self.model_path = ''
|
||||
if self.model_path == '':
|
||||
print('No correct file location has been given for saving the model, thus it won\'t be saved')
|
||||
|
||||
try:
|
||||
self.test_set_share = float(self.setting['Test set share'])
|
||||
|
||||
# wrong value
|
||||
diff = (1 - self.test_set_share)
|
||||
if diff <= 0 or diff > 1:
|
||||
self.test_set_share = 0
|
||||
print('The value of the test set share must be between 0 and 1 (1 not included), no prediction will be performed')
|
||||
|
||||
except KeyError:
|
||||
self.test_set_share = 0
|
||||
print('The value of the test set share must be between 0 and 1 (1 not included), no prediction will be performed')
|
||||
|
||||
try:
|
||||
self.save_path = self.setting['Filename to load model from']
|
||||
self.model = pickle.load(open(self.save_path, 'rb'))
|
||||
except KeyError:
|
||||
self.save_path = ''
|
||||
print('No correct location has been given to load the model from, thus a new model will be created.')
|
||||
|
||||
# if model doesn't exist we will init a new one with params from the box
|
||||
# FileNotFoundError doesn't exist in Python 2.7
|
||||
except IOError:
|
||||
print('No correct location has been given to load the model from, thus a new model will be created.')
|
||||
|
||||
# special case for Riemannian Geometry because it needs a pipeline
|
||||
clf = self.setting['Classifier']
|
||||
|
||||
try:
|
||||
discriminator, _ = self.map_clf(self.setting['Discriminator'])
|
||||
except KeyError:
|
||||
discriminator = None
|
||||
|
||||
if clf == 'Riemann Tangent Space':
|
||||
|
||||
if discriminator is not None:
|
||||
self.clf = make_pipeline(Covariances(), TangentSpace(metric='riemann'), discriminator())
|
||||
else:
|
||||
self.clf = make_pipeline(Covariances(), TangentSpace(metric='riemann'), LinearDiscriminantAnalysis())
|
||||
|
||||
elif clf == 'Riemann Minimum Distance to Mean':
|
||||
if discriminator is not None:
|
||||
self.clf = make_pipeline(Covariances(), MDM(metric=dict(mean='riemann', distance='riemann')), discriminator())
|
||||
else:
|
||||
self.clf = make_pipeline(Covariances(), MDM(metric=dict(mean='riemann', distance='riemann')))
|
||||
|
||||
else:
|
||||
self.clf, _ = self.map_clf(clf)
|
||||
self.init_params()
|
||||
|
||||
try:
|
||||
self.clf = self.clf(**self.clf_dependant_settings)
|
||||
except TypeError:
|
||||
self.clf = self.clf()
|
||||
|
||||
#----------------------------------------
|
||||
def init_params(self):
|
||||
|
||||
# default settings
|
||||
self.std_settings = ['Classifier', 'Discriminator', 'Filename to load model from', 'Test set share',
|
||||
'Filename to save configuration to', 'Filename to save model to', 'Clock frequency (Hz)', 'Labels']
|
||||
settings = [key for key in self.setting.keys()]
|
||||
|
||||
# we get only settings that are clf relevant
|
||||
clf_dependant_settings = list(set(settings) - set(self.std_settings))
|
||||
clf_dependant_settings = dict((k, v) for k, v in self.setting.items() if (k in clf_dependant_settings and len(v) > 0))
|
||||
|
||||
# we convert values that need to be
|
||||
for k, v in clf_dependant_settings.items():
|
||||
|
||||
try:
|
||||
expr = eval(v)
|
||||
clf_dependant_settings[k] = expr
|
||||
except:
|
||||
if v.lower() == 'true':
|
||||
clf_dependant_settings[k] = True
|
||||
elif v.lower() == 'false':
|
||||
clf_dependant_settings[k] = False
|
||||
elif v.lower() == 'none':
|
||||
clf_dependant_settings[k] = None
|
||||
else:
|
||||
pass
|
||||
|
||||
self.clf_dependant_settings = clf_dependant_settings
|
||||
|
||||
#----------------------------------------
|
||||
def map_clf(self, classifier):
|
||||
"""
|
||||
Returns the correct algorithm according to the classifier string
|
||||
"""
|
||||
|
||||
switcher = {
|
||||
'': None,
|
||||
'None': None,
|
||||
'Nearest Centroid': NearestCentroid,
|
||||
'Nearest Neighbors Classifier': KNeighborsClassifier,
|
||||
'Gaussian Naive Bayes': GaussianNB,
|
||||
'Stochastic Gradient Descent': SGDClassifier,
|
||||
'Logistic Regression': LogisticRegression,
|
||||
'Decision Tree Classifier': DecisionTreeClassifier,
|
||||
'Extra Trees': ExtraTreesClassifier,
|
||||
'Bagging': BaggingClassifier,
|
||||
'Random Forest': RandomForestClassifier,
|
||||
'Support Vector Machine': LinearSVC,
|
||||
'Linear Discriminant Analysis': LinearDiscriminantAnalysis,
|
||||
'AdaBoost': AdaBoostClassifier,
|
||||
'Multi Layer Perceptron': MLPClassifier,
|
||||
'Linear SVC': LinearSVC,
|
||||
}
|
||||
clf = switcher.get(classifier, lambda: 'unknown classifier')
|
||||
return clf, switcher
|
||||
|
||||
#----------------------------------------
|
||||
def on_chunk_received(self, chunk, label, shape):
|
||||
|
||||
# special case for riemannian geometry
|
||||
if self.setting['Classifier'] == 'Riemann Minimum Distance to Mean' or self.setting['Classifier'] == 'Riemann Tangent Space':
|
||||
numpyBuffer = np.array(chunk).reshape(shape)
|
||||
self.data[label][-1] = numpyBuffer
|
||||
|
||||
#----------------------------------------
|
||||
def on_end_box(self):
|
||||
try:
|
||||
self.train()
|
||||
self.save()
|
||||
except Exception as e:
|
||||
print(e)
|
||||
self.send_end_stim()
|
||||
|
||||
#----------------------------------------
|
||||
def train(self):
|
||||
x_train, y_train, x_test, y_test = train_test_split(self.data, self.test_set_share)
|
||||
|
||||
if self.model != None:
|
||||
self.clf = self.model
|
||||
|
||||
self.clf.fit(x_train, y_train)
|
||||
|
||||
# to be used in ProcessML
|
||||
self.clf.custom_classifier = self.setting['Classifier']
|
||||
|
||||
if x_test.shape[0] > 0:
|
||||
predictions = self.clf.predict(x_test)
|
||||
report = classification_report(y_test, predictions, labels=list(self.data.keys()))
|
||||
matrix = confusion_matrix(y_test, predictions)
|
||||
|
||||
print("Report :\n{}\n".format(report))
|
||||
print("Confusion Matrix : \n{}\n".format(matrix))
|
||||
print("Fin de l'entrainement...")
|
||||
|
||||
#----------------------------------------
|
||||
def save(self):
|
||||
if self.model_path != "":
|
||||
pickle.dump(self.clf, open(self.model_path, 'wb'))
|
||||
print('Model saved in {}\n'.format(self.model_path))
|
||||
|
||||
|
||||
box = TrainerML()
|
||||
Reference in New Issue
Block a user