|
| 1 | +"""Train SNN on ASL-DVS (American Sign Language) benchmark. |
| 2 | +
|
| 3 | +Architecture: 2048 -> hidden1 (recurrent adLIF) -> hidden2 (adLIF) -> 24 |
| 4 | +24-class event-camera sign language letter classification. |
| 5 | +
|
| 6 | +Usage: |
| 7 | + python asl_dvs/train.py --epochs 200 --device cuda:0 |
| 8 | +""" |
| 9 | + |
| 10 | +import os |
| 11 | +import sys |
| 12 | +import argparse |
| 13 | +import random |
| 14 | +import numpy as np |
| 15 | +import torch |
| 16 | +import torch.nn as nn |
| 17 | + |
| 18 | +sys.path.insert(0, os.path.dirname(os.path.dirname(__file__))) |
| 19 | + |
| 20 | +from common.neurons import LIFNeuron, AdaptiveLIFNeuron, surrogate_spike |
| 21 | +from common.training import run_training |
| 22 | +from common.augmentation import event_drop |
| 23 | + |
| 24 | +from asl_dvs.loader import ASLDVSDataset, collate_fn, N_CHANNELS, N_CLASSES |
| 25 | + |
| 26 | + |
| 27 | +class ASLDVSSNN(nn.Module): |
| 28 | + """Two-layer SNN for ASL-DVS classification.""" |
| 29 | + |
| 30 | + def __init__(self, n_input=N_CHANNELS, n_hidden1=512, n_hidden2=256, |
| 31 | + n_output=N_CLASSES, beta_hidden=0.95, beta_out=0.9, |
| 32 | + threshold=1.0, dropout=0.3, neuron_type='adlif', |
| 33 | + alpha_init=0.93, rho_init=0.85, beta_a_init=0.05): |
| 34 | + super().__init__() |
| 35 | + self.n_hidden1 = n_hidden1 |
| 36 | + self.n_hidden2 = n_hidden2 |
| 37 | + self.n_output = n_output |
| 38 | + self.neuron_type = neuron_type |
| 39 | + |
| 40 | + self.fc1 = nn.Linear(n_input, n_hidden1, bias=False) |
| 41 | + self.fc_rec = nn.Linear(n_hidden1, n_hidden1, bias=False) |
| 42 | + self.fc2 = nn.Linear(n_hidden1, n_hidden2, bias=False) |
| 43 | + self.fc3 = nn.Linear(n_hidden2, n_output, bias=False) |
| 44 | + |
| 45 | + if neuron_type == 'adlif': |
| 46 | + self.lif1 = AdaptiveLIFNeuron( |
| 47 | + n_hidden1, alpha_init=alpha_init, rho_init=rho_init, |
| 48 | + beta_a_init=beta_a_init, threshold=threshold) |
| 49 | + self.lif2 = AdaptiveLIFNeuron( |
| 50 | + n_hidden2, alpha_init=alpha_init, rho_init=rho_init, |
| 51 | + beta_a_init=beta_a_init, threshold=threshold) |
| 52 | + else: |
| 53 | + self.lif1 = LIFNeuron(n_hidden1, beta_init=beta_hidden, |
| 54 | + threshold=threshold, learn_beta=True) |
| 55 | + self.lif2 = LIFNeuron(n_hidden2, beta_init=beta_hidden, |
| 56 | + threshold=threshold, learn_beta=True) |
| 57 | + |
| 58 | + self.lif_out = LIFNeuron(n_output, beta_init=beta_out, |
| 59 | + threshold=threshold, learn_beta=True) |
| 60 | + self.dropout1 = nn.Dropout(p=dropout) |
| 61 | + self.dropout2 = nn.Dropout(p=dropout) |
| 62 | + |
| 63 | + nn.init.xavier_uniform_(self.fc1.weight, gain=0.5) |
| 64 | + nn.init.xavier_uniform_(self.fc2.weight, gain=0.5) |
| 65 | + nn.init.xavier_uniform_(self.fc3.weight, gain=0.5) |
| 66 | + nn.init.orthogonal_(self.fc_rec.weight, gain=0.2) |
| 67 | + |
| 68 | + def forward(self, x): |
| 69 | + batch, T, _ = x.shape |
| 70 | + device = x.device |
| 71 | + |
| 72 | + v1 = torch.zeros(batch, self.n_hidden1, device=device) |
| 73 | + v2 = torch.zeros(batch, self.n_hidden2, device=device) |
| 74 | + v_out = torch.zeros(batch, self.n_output, device=device) |
| 75 | + spk1 = torch.zeros(batch, self.n_hidden1, device=device) |
| 76 | + spk2 = torch.zeros(batch, self.n_hidden2, device=device) |
| 77 | + out_sum = torch.zeros(batch, self.n_output, device=device) |
| 78 | + |
| 79 | + if self.neuron_type == 'adlif': |
| 80 | + a1 = torch.zeros(batch, self.n_hidden1, device=device) |
| 81 | + a2 = torch.zeros(batch, self.n_hidden2, device=device) |
| 82 | + |
| 83 | + for t in range(T): |
| 84 | + I1 = self.fc1(x[:, t]) + self.fc_rec(spk1) |
| 85 | + if self.neuron_type == 'adlif': |
| 86 | + v1, spk1, a1 = self.lif1(I1, v1, a1, spk1) |
| 87 | + else: |
| 88 | + v1, spk1 = self.lif1(I1, v1) |
| 89 | + spk1_d = self.dropout1(spk1) if self.training else spk1 |
| 90 | + |
| 91 | + I2 = self.fc2(spk1_d) |
| 92 | + if self.neuron_type == 'adlif': |
| 93 | + v2, spk2, a2 = self.lif2(I2, v2, a2, spk2) |
| 94 | + else: |
| 95 | + v2, spk2 = self.lif2(I2, v2) |
| 96 | + spk2_d = self.dropout2(spk2) if self.training else spk2 |
| 97 | + |
| 98 | + I_out = self.fc3(spk2_d) |
| 99 | + beta_out = self.lif_out.beta |
| 100 | + v_out = beta_out * v_out + (1.0 - beta_out) * I_out |
| 101 | + out_sum = out_sum + v_out |
| 102 | + |
| 103 | + return out_sum / T |
| 104 | + |
| 105 | + |
| 106 | +def main(): |
| 107 | + parser = argparse.ArgumentParser(description="Train SNN on ASL-DVS") |
| 108 | + parser.add_argument("--data-dir", default="data/asl_dvs") |
| 109 | + parser.add_argument("--epochs", type=int, default=200) |
| 110 | + parser.add_argument("--batch-size", type=int, default=64) |
| 111 | + parser.add_argument("--lr", type=float, default=1e-3) |
| 112 | + parser.add_argument("--weight-decay", type=float, default=1e-4) |
| 113 | + parser.add_argument("--hidden1", type=int, default=512) |
| 114 | + parser.add_argument("--hidden2", type=int, default=256) |
| 115 | + parser.add_argument("--dropout", type=float, default=0.3) |
| 116 | + parser.add_argument("--time-bins", type=int, default=10) |
| 117 | + parser.add_argument("--seed", type=int, default=42) |
| 118 | + parser.add_argument("--save", default="asl_dvs_model.pt") |
| 119 | + parser.add_argument("--neuron", choices=["lif", "adlif"], default="adlif") |
| 120 | + parser.add_argument("--alpha-init", type=float, default=0.93) |
| 121 | + parser.add_argument("--rho-init", type=float, default=0.85) |
| 122 | + parser.add_argument("--beta-a-init", type=float, default=0.05) |
| 123 | + parser.add_argument("--event-drop", action="store_true", default=True) |
| 124 | + parser.add_argument("--label-smoothing", type=float, default=0.05) |
| 125 | + parser.add_argument("--device", default=None) |
| 126 | + args = parser.parse_args() |
| 127 | + |
| 128 | + torch.manual_seed(args.seed) |
| 129 | + np.random.seed(args.seed) |
| 130 | + random.seed(args.seed) |
| 131 | + |
| 132 | + device = torch.device(args.device or ("cuda" if torch.cuda.is_available() else "cpu")) |
| 133 | + print(f"Device: {device}") |
| 134 | + |
| 135 | + print("Loading ASL-DVS dataset...") |
| 136 | + train_ds = ASLDVSDataset(args.data_dir, train=True, n_time_bins=args.time_bins) |
| 137 | + test_ds = ASLDVSDataset(args.data_dir, train=False, n_time_bins=args.time_bins) |
| 138 | + |
| 139 | + from torch.utils.data import DataLoader |
| 140 | + train_loader = DataLoader( |
| 141 | + train_ds, batch_size=args.batch_size, shuffle=True, |
| 142 | + collate_fn=collate_fn, num_workers=0, pin_memory=True) |
| 143 | + test_loader = DataLoader( |
| 144 | + test_ds, batch_size=args.batch_size, shuffle=False, |
| 145 | + collate_fn=collate_fn, num_workers=0, pin_memory=True) |
| 146 | + |
| 147 | + print(f"Train: {len(train_ds)}, Test: {len(test_ds)}, Time bins: {args.time_bins}") |
| 148 | + |
| 149 | + model = ASLDVSSNN( |
| 150 | + n_hidden1=args.hidden1, n_hidden2=args.hidden2, |
| 151 | + dropout=args.dropout, neuron_type=args.neuron, |
| 152 | + alpha_init=args.alpha_init, rho_init=args.rho_init, |
| 153 | + beta_a_init=args.beta_a_init, |
| 154 | + ).to(device) |
| 155 | + |
| 156 | + print(f"Model: {N_CHANNELS}->{args.hidden1}->{args.hidden2}->{N_CLASSES} " |
| 157 | + f"({args.neuron.upper()}, recurrent=on, dropout={args.dropout})") |
| 158 | + |
| 159 | + augment_fn = (lambda x: event_drop(x)) if args.event_drop else None |
| 160 | + |
| 161 | + config = { |
| 162 | + 'device': device, 'epochs': args.epochs, 'lr': args.lr, |
| 163 | + 'weight_decay': args.weight_decay, 'save_path': args.save, |
| 164 | + 'benchmark': 'asl_dvs', 'augment_fn': augment_fn, |
| 165 | + 'label_smoothing': args.label_smoothing, |
| 166 | + 'model_config': { |
| 167 | + 'n_input': N_CHANNELS, 'hidden1': args.hidden1, |
| 168 | + 'hidden2': args.hidden2, 'n_output': N_CLASSES, |
| 169 | + 'neuron_type': args.neuron, 'dropout': args.dropout, |
| 170 | + }, |
| 171 | + } |
| 172 | + run_training(model, train_loader, test_loader, config) |
| 173 | + |
| 174 | + |
| 175 | +if __name__ == "__main__": |
| 176 | + main() |
0 commit comments