Tarea 5: Clasificador con redes de picos

Tarea 5: Clasificador con redes de picos

Información de la Tarea


Descripción de la Tarea

Análisis del Conjunto de Datos El Reto

Muestra: 4,898 vinos y 11 características químicas.

El Desbalance:

Distribución de clases

Es fundamental destacar que el ~75% de los datos se concentran en las calidades 5 y 6. Esto define la barrera matemática del problema, ya que calidades extremas (3 o 9) representan menos del 1% de la muestra.

Diseño de la Arquitectura SNN

Hiperparámetros Clave:

import torch
import torch.nn as nn
import snntorch as snn
class Net(nn.Module):
def __init__(self, num_inputs, num_hidden, num_outputs, beta, num_steps):
super().__init__()
self.num_inputs = num_inputs
self.num_hidden = num_hidden
self.num_outputs = num_outputs
self.beta = beta
self.num_steps = num_steps
self.InputMatrix = nn.Linear(self.num_inputs, self.num_hidden, False)
self.lif1 = snn.RLeaky(beta=beta, linear_features=num_hidden)
self.lif1.recurrent.bias.requires_grad = False
self.lif1.recurrent.bias.fill_(0)
self.OutputMatrix = nn.Linear(self.num_hidden, self.num_outputs, False)
self.lif2 = snn.Leaky(beta=beta)
def forward(self, x):
spk1, mem1 = self.lif1.init_rleaky()
mem2 = self.lif2.init_leaky()
spk2_out = []
for step in range(self.num_steps):
InputPropagation = self.InputMatrix(x)
spk1, mem1 = self.lif1(InputPropagation, spk1, mem1)
OutputPropagation = self.OutputMatrix(spk1)
spk2, mem2 = self.lif2(OutputPropagation, mem2)
spk2_out.append(spk2)
return torch.stack(spk2_out, dim=0)

Proceso de Entrenamiento PyTorch

Configuración e Hiperparámetros:

Resultados de Entrenamiento:

Curva de Pérdida

La pérdida (CrossEntropyLoss) disminuye progresivamente, demostrando la convergencia de la red de picos.

Evaluación

Una vez entrenada la red, se extrajeron las 3 matrices de pesos absolutas ($C_2, C_3, C_4$) para realizar la inferencia y evaluar matemáticamente el modelo con NumPy:

import numpy as np
C2 = np.array([
[-2.39488053, -3.16118383, -2.85059690, -3.02237344, -0.41534084, -2.92052913, -1.80023563, -2.57450676, -3.32677388, -1.51347661, -4.47367287],
[-6.11782122, -5.13261652, -6.38331652, -3.78472662, -5.00517464, -5.69930935, -5.99944592, -8.32645130, -6.53935957, -5.33293772, -6.63186550],
[0.47290173, -0.33224255, -2.05781722, -0.35796332, 3.14016485, 1.27089560, 1.88768625, 0.54269075, -0.60878360, 0.04172051, 0.11060597],
[-1.52368045, 1.84668839, -0.59144539, 0.39053944, 2.06824517, 0.84178454, 0.29720491, 1.26986802, 0.80111784, -0.40023351, -0.37118641],
[-5.13319016, -6.39248276, -6.32654810, -1.28760099, -2.99321747, -5.56891537, -4.12379646, -4.18455172, -7.07475281, -6.12247276, -7.73189211],
[0.98943663, 0.29252395, 0.96952665, 0.25094491, -0.03749528, -2.00045919, -0.34092790, 0.49472892, 0.51288515, 0.72585976, 0.79764020],
[-1.31664944, -1.20742226, -0.49209425, 1.91920435, 1.93731558, 0.82496035, 0.27434182, -1.48550248, -1.97232246, 1.23482144, 0.36456063],
[-1.43486929, -1.74014616, -2.87549782, -1.99650097, -2.35450315, 0.02938510, -0.65274715, -2.08736038, -1.73759270, -1.25210679, -3.94077301],
[-1.00994575, -0.98603129, -0.35445717, 0.90468931, 0.40140504, 4.01333714, 0.14599289, -1.22081900, 0.44824323, 0.13372605, 0.89749980],
[0.22962669, -1.84435427, -0.05217035, 1.22189963, -0.39794368, -2.14430475, 0.75829703, 0.78106350, -0.15984923, 0.89997363, 0.56958908],
[1.48051214, -1.55657780, 0.21288587, 0.39003596, -1.39649189, 0.05906558, -1.10666692, -0.10105289, 0.53379577, 1.22400260, -2.42923713],
[0.56534868, 0.14357370, -0.26594511, 1.80443943, -1.02878511, 1.01654851, 0.09173497, -3.43639040, 0.55007732, 0.22668192, -0.00310216],
[0.42732498, 1.85627961, 0.66570216, -0.97733587, 1.08376646, -0.67432231, 1.00610495, 0.83892739, -0.28163108, -0.03900059, -1.10337365],
[-8.13457394, -8.31346226, -8.27906609, -3.35346031, -7.09960890, -7.44434404, -6.64564753, -9.38786221, -7.73659086, -7.18230677, -7.41583443],
[-3.91701341, -7.06561089, -3.97988129, 0.95131314, -2.24017477, -5.27675581, -4.76773357, -4.59004736, -5.11132050, -3.26487660, -4.32863379],
[-0.28016907, 0.87928730, -0.69748980, -3.22336054, -1.33296406, -3.31045961, -1.14311814, 2.48456979, 0.81482339, -0.20544863, 0.51130432],
[-0.80769718, 1.01975179, -0.19339615, -0.17835237, 1.17013097, 1.97420537, -0.67243874, 1.16255283, -0.57065195, -2.29249239, -0.16528387],
[-0.79824609, 0.80101949, -0.99181503, -0.03243065, 0.43425000, -1.46199870, 0.29875806, 1.02606463, 0.45220459, 0.62988502, -0.84390360],
[1.07260442, -2.74451065, -0.83007801, 0.45637280, 1.89356184, -1.45044374, 0.59085470, 0.50102562, -0.09532350, -0.59341031, -1.09225571],
[-0.84504610, -2.64524317, -1.39739621, 0.02565018, 0.28001782, -2.17073083, -1.11969149, -1.35730410, -1.76431227, -1.59121323, -2.82267714],
])
C3 = np.array([
[-0.59411627, -0.22192989, -0.83276296, -0.85958713, 0.05652140, -0.89175069, -1.10807562, -0.04273168, -1.84221029, 0.30033448, 1.28147149, 0.84465694, -1.35121942, 0.17129309, 0.05016008, -0.94370109, 0.50436777, -1.05249059, -1.12154138, 0.19208580],
[-0.06417506, 0.10550968, -2.10051203, -1.29400611, 0.09926854, -3.83859658, 2.42378235, -0.06017752, -5.46468067, -0.24243584, 1.04112923, -0.37156001, -5.82276058, -0.05047070, -0.18679203, -1.46629715, -1.30354595, -2.42850614, 1.46589267, -0.16097927],
[0.22515270, 0.02375370, -1.89384401, -0.12230647, 0.21610263, -1.66464710, -0.09651586, -0.31626227, -0.19506420, -0.59810519, -0.08286761, 0.04144407, 0.02605123, 0.17782788, -0.03224108, -1.73499691, -0.13564363, -0.29998907, -1.33805466, 0.17483927],
[-0.31262109, 0.16398661, -0.40979713, -0.09927863, 0.07211636, -0.83723396, -1.22662318, 0.22769529, 0.12660117, 0.43774760, 1.64660299, -0.62003702, -1.64178813, 0.05837956, 0.15893775, -3.53080177, -1.20706785, 0.38802108, -0.71367586, -0.21323434],
[0.20786744, 0.20875281, -0.80062878, -2.13488817, -0.02171124, -5.31066179, -0.14277501, 0.08717147, -6.54711294, 1.32438326, 0.44965175, -2.28111792, -6.21658659, 0.06288467, -0.14103091, -0.96143150, -0.40776122, -3.25709200, 2.16652203, -0.14301196],
[-0.32896504, -0.03312289, -0.90666729, 1.02320790, -0.20103668, -0.83674598, 1.05355597, 0.96604377, -0.66831416, 0.53937113, -1.11980927, -1.73505640, -0.30766717, 0.13917293, 0.10399097, 0.14040941, 2.09063888, 0.14408261, -1.06440258, 0.06211904],
[0.30102572, 0.05138361, 0.91332054, -0.10183641, -0.06500120, -1.38394332, 0.57999212, -0.70403939, 0.67589992, -0.27090612, 1.41017342, 0.66778362, -1.22596037, -0.01344702, -0.17049481, 0.96574646, 0.16208188, -4.85798168, 0.82143945, -0.06004349],
[-0.08531620, -0.18575263, -0.80953729, -0.37945518, 0.05462424, -0.45692858, 0.56521803, 0.74951142, -1.39847231, -1.51909637, 2.22999644, 2.29062486, 2.04797459, 0.10210965, 0.05543762, 3.26581693, -1.49815917, 1.17570627, -1.02161610, 0.02397856],
[-0.23024169, 0.04915357, 1.06514287, 0.78185296, -0.08454566, -0.56864321, -1.29345298, 0.25062701, 0.57254100, 0.28046834, 0.10606694, 0.45142114, 0.29085571, 0.01963540, 0.15668708, -0.21991926, -0.81473911, -0.01061889, 1.10626876, 0.24036269],
[-0.42813089, -0.11441419, 1.79262042, -3.15963387, 0.14871530, -1.76148999, 0.09149078, -0.19203915, -0.29325694, 0.12443326, -1.32197189, 0.26916027, 0.30210641, -0.05205239, 0.21307248, 1.20274007, 1.03454423, 0.84049249, -0.28827888, 0.04049236],
[0.14646065, 0.00126811, -2.48267841, -0.41298342, 0.12138522, 0.10950031, -4.17488670, 0.44834605, -0.56547374, -0.51251948, 1.21395791, 0.18426940, -0.63238424, -0.13790080, 0.02153661, 2.15245032, -0.70717239, 0.08708605, 0.21299709, 0.16364777],
[0.29710257, 0.15558602, -0.66139674, 0.49967498, 0.14836419, -0.07653025, 0.37196085, -1.00090384, -0.27698761, -0.52519065, 0.62742889, 0.85496056, 0.40723827, 0.22166683, 0.19360949, -0.21290210, 0.66190791, -0.61233252, -1.05246806, 0.22119361],
[-0.06770706, 0.11276411, -0.26105961, 0.10231667, 0.01794293, 1.27118123, 0.44635436, 0.75579286, -0.44916725, -0.98175186, -0.15365833, -0.27987736, 0.50833738, -0.14210448, -0.13573718, 0.92064506, 0.51289177, 0.33564717, -0.47506008, 0.19704397],
[-0.09152219, 0.03266222, -2.95632601, -1.66078424, 0.15666234, -5.22585726, 1.06370950, 0.01172664, -6.69583464, 0.07171202, 0.50245190, -0.68250334, -6.50504065, 0.07772394, 0.11444274, -2.63167667, -1.49495912, -2.66794515, 0.10763498, -0.21053112],
[0.16542445, 0.16384925, -0.67645776, -2.29077101, 0.00351563, -2.79754591, -0.00338659, 0.24387334, -3.99311376, 1.79548633, 2.55405736, 1.11809587, -4.27903605, -0.20610087, 0.07279137, -2.38863683, -0.79696184, -2.53205967, 1.17150950, 0.07162819],
[0.07610549, -0.16569643, 2.10868549, -0.60723668, 0.01781290, 0.70359635, -1.83066809, -0.79234397, -1.17556119, 0.38311809, 1.12588668, 0.01386364, -0.28024301, 0.11576907, 0.20493418, 0.01480751, 0.46309930, 0.18745263, 2.21869826, -0.19145879],
[-0.14969325, 0.16115268, 2.17959881, 0.57308984, 0.05603223, 0.10231651, 0.85683227, -0.64032823, 0.25635573, -0.02655266, -3.72723126, 0.70714206, -0.22615638, -0.14020300, -0.01964801, 1.24455500, -1.68225586, 0.75307834, 0.06969352, -0.01277610],
[0.21331561, -0.09833988, 0.62585318, 1.45441413, 0.02189370, -0.71524078, -0.02659469, -0.11509438, 0.46844780, -0.33219406, 0.08218818, -0.31337640, 0.20745111, 0.15670891, 0.18422025, -0.13465174, 0.55603266, -1.90883493, -0.75341290, 0.20115426],
[0.23976815, -0.04694869, 1.17970288, 0.27066755, 0.16913886, -0.53907174, -0.66834408, -0.41638207, -0.09093174, -0.19094872, 1.49376214, -1.30298936, 0.41747403, 0.05217849, 0.07487334, -0.01289409, 0.07166098, 0.42162359, 1.13319933, 0.19435890],
[-0.10111918, 0.10992340, 0.29692143, 0.74695587, -0.19036362, -0.05430766, -0.72498286, -0.82101101, -3.76743817, 0.62149853, 0.17795132, 1.21118081, 1.07319593, 0.18661241, -0.11808398, -0.49146158, 2.60854530, 0.39558053, 3.20398092, 0.02929873],
])
C4 = np.array([
[-0.36661875, -0.13097605, 1.34947681, -1.20681834, -0.11157879, 0.07779110, -1.08558691, -1.05536830, -1.35023808, -9.40265274, -0.22081937, -0.92387140, 0.65452671, -0.20301607, 0.03754491, 0.02668205, -0.65696657, 0.49472907, 1.06572247, -0.20775841],
[0.08499817, -0.01959275, -0.30732265, -0.79475409, 0.14301716, 0.44694102, 0.12647907, -0.23851955, 0.01932184, -1.22356868, -0.18933620, -0.26022112, 0.73137027, 0.18996094, 0.15059514, -0.01257388, -0.45771587, 0.52695322, -0.10198086, -0.29747596],
[-0.04230051, -0.16379759, -0.18500902, 0.08662852, -0.20589729, 0.48552385, -0.82535821, 0.38789117, 0.41359201, -0.24746135, 0.03814263, -0.32207602, 0.89253294, -0.13420479, -0.22304833, -0.63776559, -0.64824855, -0.13205542, -0.02551257, -0.10482457],
[-0.07321698, -0.11708526, -0.44255593, -0.25758639, -0.04168928, 0.97284091, -0.45134717, -0.95407110, 0.62571448, -0.20972973, 0.28724325, -0.21269965, 0.53114623, -0.12502813, -0.00165894, -0.70455068, -0.76463902, -0.18074015, 0.07329077, -0.03389637],
[-0.27458802, 0.09292007, -0.61633772, -0.54780805, -0.00159891, 0.97652233, 0.14449230, -1.53581607, 0.61388737, -0.32048610, -0.14684872, 0.07234244, 0.06986043, -0.03177399, 0.04526694, -0.51838005, -2.24674535, -0.19643307, 0.40019378, 0.01358079],
[-0.41662610, 0.06537189, -2.11232185, -0.67029840, 0.19397512, 0.84308678, 0.69671869, -1.79099226, 0.43884346, -0.58899796, -0.27551261, 0.12003243, -0.09126518, 0.04092953, 0.01436823, -0.50818604, -0.42884105, -1.11883521, 0.51684314, 0.06445422],
[-0.39658469, 0.11978407, -27.01889992, -3.26441264, 0.02746528, 1.00007439, -20.57846832, -0.87039906, -2.11781120, -21.23101234, -10.01791191, 2.85906577, -5.52525043, 0.21446453, -0.02141666, -15.91802883, -2.70464468, -29.95018005, -21.16488266, 0.07955072],
])
def leaky(valin, spk, mem, beta):
mem = beta * mem + valin - spk
if mem > 1.0:
spk = 1.0
else:
spk = 0.0
return spk, mem
def rleaky(valin, recurrent, spk, mem, beta):
mem = beta * mem + valin + recurrent - spk
if mem > 1.0:
spk = 1.0
else:
spk = 0.0
return spk, mem
def evaluate(vin, steps=10):
spk1 = np.zeros(20)
mem1 = np.zeros(20)
spk2 = np.zeros(7)
mem2 = np.zeros(7)
spk2_out = []
spk1_out = []
beta = 0.95
for i in range(steps):
vs2 = C2 @ vin
vs3 = C3 @ spk1
for j in range(20):
spk1[j], mem1[j] = rleaky(vs2[j], vs3[j], spk1[j], mem1[j], beta)
vs4 = C4 @ spk1
for j in range(7):
spk2[j], mem2[j] = leaky(vs4[j], spk2[j], mem2[j], beta)
spk2_out.append(spk2.copy())
spk1_out.append(spk1.copy())
return np.stack(spk2_out), np.stack(spk1_out)

Comportamiento de Picos:

Raster Plot - Picos de Prueba

Se muestra cómo se comunican las neuronas ocultas con la capa de salida a lo largo de los 10 pasos para emitir el veredicto final.

Resultados Finales y Matriz de Confusión

Reporte de Entrenamiento:

Matriz de Confusión

Diagnóstico: El modelo separa eficientemente los vinos promedio (5, 6 y 7). Los errores en clases extremas (3 y 9) son resultado del alto desbalance, por lo que la red optimiza apostando a las categorías con mayor densidad de datos.