diff --git a/.gitignore b/.gitignore index 9713d0a..625f716 100644 --- a/.gitignore +++ b/.gitignore @@ -16,4 +16,5 @@ sequences cifar-data/ .DS_Store *.parquet -.mypy_cache/ \ No newline at end of file +.mypy_cache/ +*.log \ No newline at end of file diff --git a/RecurrentFF/benchmarks/mnist/mnist.py b/RecurrentFF/benchmarks/mnist/mnist.py index aaabd80..90e7754 100644 --- a/RecurrentFF/benchmarks/mnist/mnist.py +++ b/RecurrentFF/benchmarks/mnist/mnist.py @@ -12,7 +12,7 @@ DATA_SIZE = 784 NUM_CLASSES = 10 TRAIN_BATCH_SIZE = 500 -TEST_BATCH_SIZE = 5000 +TEST_BATCH_SIZE = 500 ITERATIONS = 15 DATASET = "MNIST" diff --git a/RecurrentFF/model/data_scenario/static_single_class.py b/RecurrentFF/model/data_scenario/static_single_class.py index ed57166..62834f4 100644 --- a/RecurrentFF/model/data_scenario/static_single_class.py +++ b/RecurrentFF/model/data_scenario/static_single_class.py @@ -1,3 +1,4 @@ +import copy import logging from typing import List, Optional, cast from pyparsing import Iterator @@ -378,111 +379,124 @@ class label, using a two-step process: and is_test_set, "Cannot write activations for batch size > 1" activity_tracker = StaticSingleClassActivityTracker() - forward_mode = ForwardMode.PredictData if is_test_set else ForwardMode.PositiveData + # forward_mode = ForwardMode.PredictData if is_test_set else ForwardMode.PositiveData + forward_mode = ForwardMode.PositiveData # tuple: (correct, total) accuracy_contexts = [] + # save acts + activations_saved = [] + for layer in self.inner_layers: + layer_acts = copy.deepcopy(layer.pos_activations) + activations_saved.append(layer_acts) + torch.save(self.inner_layers.state_dict(), "./tmp_model.pt") + for batch, test_data in enumerate(loader): if limit_batches is not None and batch == limit_batches: break - with torch.no_grad(): - data, labels = test_data - data = data.to(self.settings.device.device) - labels = labels.to(self.settings.device.device) + data, labels = test_data + data = data.to(self.settings.device.device) + labels = labels.to(self.settings.device.device) - if write_activations: - activity_tracker.reinitialize(data, labels) - - # since this is static singleclass we can use the first frame - # for the label - labels = labels[0] - - iterations = data.shape[0] - - all_labels_badness = [] - - # evaluate badness for each possible label - for label in range(self.settings.data_config.num_classes): - self.inner_layers.reset_activations(not is_test_set) - - upper_clamped_tensor = self.get_preinit_upper_clamped_tensor( - (data.shape[1], self.settings.data_config.num_classes)) - - for _preinit_step in range( - 0, self.settings.model.prelabel_timesteps): - self.inner_layers.advance_layers_forward( - forward_mode, data[0], upper_clamped_tensor, False) - if write_activations: - activity_tracker.track_partial_activations( - self.inner_layers) - - one_hot_labels = torch.zeros( - data.shape[1], - self.settings.data_config.num_classes, - device=self.settings.device.device) - one_hot_labels[:, label] = 1.0 - - lower_iteration_threshold = iterations // 2 - \ - iterations // 10 - upper_iteration_threshold = iterations // 2 + \ - iterations // 10 - badnesses = [] - for iteration in range(0, iterations): - self.inner_layers.advance_layers_forward( - forward_mode, data[iteration], one_hot_labels, True) - if write_activations: - activity_tracker.track_partial_activations( - self.inner_layers) - - if iteration >= lower_iteration_threshold and iteration <= upper_iteration_threshold: - layer_badnesses = [] - for layer in self.inner_layers: - activations = cast(Activations, layer.pos_activations).current \ - if forward_mode == ForwardMode.PositiveData \ - else cast(Activations, layer.predict_activations).current - - layer_badnesses.append( - layer_activations_to_badness( - activations)) - - badnesses.append(torch.stack( - layer_badnesses, dim=1)) + if write_activations: + activity_tracker.reinitialize(data, labels) + + # since this is static singleclass we can use the first frame + # for the label + labels = labels[0] + + iterations = data.shape[0] + + all_labels_badness = [] + + # evaluate badness for each possible label + for label in range(self.settings.data_config.num_classes): + # self.inner_layers.reset_activations(True) + + # load acts + self.inner_layers.load_state_dict(torch.load( + "./tmp_model.pt", map_location=self.settings.device.device)) + for i, layer in enumerate(self.inner_layers): + layer.pos_activations = copy.deepcopy(activations_saved[i]) + upper_clamped_tensor = self.get_preinit_upper_clamped_tensor( + (data.shape[1], self.settings.data_config.num_classes)) + + for _preinit_step in range( + 0, self.settings.model.prelabel_timesteps): + self.inner_layers.advance_layers_forward( + data[0], upper_clamped_tensor, False) + if write_activations: + activity_tracker.track_partial_activations( + self.inner_layers) + + one_hot_labels = torch.zeros( + data.shape[1], + self.settings.data_config.num_classes, + device=self.settings.device.device) + one_hot_labels[:, label] = 1.0 + + lower_iteration_threshold = iterations // 2 - \ + iterations // 10 + upper_iteration_threshold = iterations // 2 + \ + iterations // 10 + badnesses = [] + for iteration in range(0, iterations): + self.inner_layers.advance_layers_forward( + data[iteration], one_hot_labels, True) if write_activations: - activity_tracker.cut_activations() + activity_tracker.track_partial_activations( + self.inner_layers) - # tensor of shape (batch_size, iterations, num_layers) - badnesses_stacked = torch.stack(badnesses, dim=1) - badnesses_mean_over_iterations = badnesses_stacked.mean( - dim=1) - badness_mean_over_layers = badnesses_mean_over_iterations.mean( - dim=1) + if iteration >= lower_iteration_threshold and iteration <= upper_iteration_threshold: + layer_badnesses = [] + for layer in self.inner_layers: + activations = cast(Activations, layer.pos_activations).current \ + if forward_mode == ForwardMode.PositiveData \ + else cast(Activations, layer.predict_activations).current - logging.debug("Badness for prediction" + " " + - str(label) + ": " + str(badness_mean_over_layers)) - all_labels_badness.append(badness_mean_over_layers) + layer_badnesses.append( + layer_activations_to_badness( + activations)) - all_labels_badness_stacked = torch.stack( - all_labels_badness, dim=1) + badnesses.append(torch.stack( + layer_badnesses, dim=1)) - # select the label with the maximum badness - predicted_labels = torch.argmin( - all_labels_badness_stacked, dim=1) if write_activations: - anti_predictions = torch.argmax( - all_labels_badness_stacked, dim=1) - activity_tracker.filter_and_persist( - predicted_labels, anti_predictions, labels) + activity_tracker.cut_activations() + + # tensor of shape (batch_size, iterations, num_layers) + badnesses_stacked = torch.stack(badnesses, dim=1) + badnesses_mean_over_iterations = badnesses_stacked.mean( + dim=1) + badness_mean_over_layers = badnesses_mean_over_iterations.mean( + dim=1) + + logging.debug("Badness for prediction" + " " + + str(label) + ": " + str(badness_mean_over_layers)) + all_labels_badness.append(badness_mean_over_layers) + + all_labels_badness_stacked = torch.stack( + all_labels_badness, dim=1) + + # select the label with the maximum badness + predicted_labels = torch.argmin( + all_labels_badness_stacked, dim=1) + if write_activations: + anti_predictions = torch.argmax( + all_labels_badness_stacked, dim=1) + activity_tracker.filter_and_persist( + predicted_labels, anti_predictions, labels) - logging.debug("Predicted labels: " + str(predicted_labels)) - logging.debug("Actual labels: " + str(labels)) + logging.debug("Predicted labels: " + str(predicted_labels)) + logging.debug("Actual labels: " + str(labels)) - total = data.size(1) - correct = (predicted_labels == labels).sum().item() + total = data.size(1) + correct = (predicted_labels == labels).sum().item() - accuracy_contexts.append((correct, total)) + accuracy_contexts.append((correct, total)) total_correct = sum(correct for correct, _total in accuracy_contexts) total_submissions = sum( diff --git a/RecurrentFF/model/hidden_layer.py b/RecurrentFF/model/hidden_layer.py index 277a6eb..3857a13 100644 --- a/RecurrentFF/model/hidden_layer.py +++ b/RecurrentFF/model/hidden_layer.py @@ -1,5 +1,7 @@ +from dataclasses import make_dataclass from enum import Enum import math +import random from typing import Dict, Optional, cast from typing_extensions import Self @@ -10,6 +12,8 @@ from torch.optim import RMSprop, Adam, Adadelta, Optimizer from torch.optim.lr_scheduler import StepLR from profilehooks import profile +from torchviz import make_dot + from RecurrentFF.util import ( Activations, @@ -65,13 +69,13 @@ def eval(self: Self) -> Self: def forward(self, mode: ForwardMode) -> torch.Tensor: if mode == ForwardMode.PositiveData: assert self.source.pos_activations is not None - source_activations = self.source.pos_activations.previous.detach() + source_activations = self.source.pos_activations.previous elif mode == ForwardMode.NegativeData: assert self.source.neg_activations is not None - source_activations = self.source.neg_activations.previous.detach() + source_activations = self.source.neg_activations.previous elif mode == ForwardMode.PredictData: assert self.source.predict_activations is not None - source_activations = self.source.predict_activations.previous.detach() + source_activations = self.source.predict_activations.previous source_activations_stdized = standardize_layer_activations( source_activations, self.source.settings.model.epsilon) @@ -275,6 +279,11 @@ def init_residual_connection(self, residual_connection: ResidualConnection) -> N def init_optimizer(self) -> None: self.optimizer: Optimizer + # param_optimizer = [(name, param) for name, param in self.named_parameters()] + # for (name, param) in param_optimizer: + # print(name) + # print(param.shape) + # input() if self.settings.model.ff_optimizer == "adam": self.optimizer = Adam(self.parameters(), lr=self.settings.model.ff_adam.learning_rate) @@ -357,44 +366,74 @@ def step_learning_rate(self) -> None: self.scheduler.step() def reset_activations(self, isTraining: bool) -> None: - activations_dim = None - if isTraining: - activations_dim = self.train_activations_dim - - pos_activations_current = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - pos_activations_previous = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - self.pos_activations = Activations( - pos_activations_current, pos_activations_previous) - - neg_activations_current = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - neg_activations_previous = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - self.neg_activations = Activations( - neg_activations_current, neg_activations_previous) - - self.predict_activations = None - - else: - activations_dim = self.test_activations_dim - - predict_activations_current = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - predict_activations_previous = torch.zeros( - activations_dim[0], activations_dim[1]).to( - self.settings.device.device) - self.predict_activations = Activations( - predict_activations_current, predict_activations_previous) - - self.pos_activations = None - self.neg_activations = None + activations_dim = self.train_activations_dim + + pos_activations_current = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + pos_activations_previous = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + self.pos_activations = Activations( + pos_activations_current, pos_activations_previous) + + neg_activations_current = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + neg_activations_previous = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + self.neg_activations = Activations( + neg_activations_current, neg_activations_previous) + + activations_dim = self.test_activations_dim + + predict_activations_current = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + predict_activations_previous = torch.zeros( + activations_dim[0], activations_dim[1], requires_grad=False).to( + self.settings.device.device) + self.predict_activations = Activations( + predict_activations_current, predict_activations_previous) + + # if isTraining: + # activations_dim = self.train_activations_dim + + # pos_activations_current = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # pos_activations_previous = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # self.pos_activations = Activations( + # pos_activations_current, pos_activations_previous) + + # neg_activations_current = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # neg_activations_previous = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # self.neg_activations = Activations( + # neg_activations_current, neg_activations_previous) + + # self.predict_activations = None + + # else: + # activations_dim = self.test_activations_dim + + # predict_activations_current = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # predict_activations_previous = torch.zeros( + # activations_dim[0], activations_dim[1], requires_grad=False).to( + # self.settings.device.device) + # self.predict_activations = Activations( + # predict_activations_current, predict_activations_previous) + + # self.pos_activations = None + # self.neg_activations = None def advance_stored_activations(self) -> None: if self.pos_activations is not None: @@ -412,8 +451,66 @@ def set_previous_layer(self, previous_layer: Self) -> None: def set_next_layer(self, next_layer: Self) -> None: self.next_layer = next_layer + def generate_lpl_loss_predictive(self, current_activations_with_grad: torch.Tensor) -> Tensor: + def generate_loss(current_act: Tensor, previous_act: Tensor) -> Tensor: + loss = (current_act - previous_act) ** 2 + loss = torch.sum(loss, dim=1) + loss = torch.sum(loss, dim=0) + loss = loss / \ + (2 * current_act.shape[0] * current_act.shape[1]) + return loss + + assert current_activations_with_grad.requires_grad == True + assert self.pos_activations.previous.requires_grad == False + pos_loss = generate_loss( + current_activations_with_grad, self.pos_activations.previous) + return pos_loss + + def generate_lpl_loss_hebbian(self, current_activations_with_grad: torch.Tensor) -> Tensor: + def generate_loss(activations: Tensor) -> Tensor: + mean_act = torch.mean(activations, dim=0) + mean_subtracted = activations - mean_act + + sigma_squared = torch.sum( + mean_subtracted ** 2, dim=0) / (activations.shape[0] - 1) + + loss = -torch.log(sigma_squared + 0.00000001).sum() / sigma_squared.shape[0] + return loss + + assert current_activations_with_grad.requires_grad == True + pos_loss = generate_loss(current_activations_with_grad) + return pos_loss + + def generate_lpl_loss_decorrelative(self, current_activations_with_grad: torch.Tensor) -> Tensor: + def generate_loss(activations: torch.Tensor) -> torch.Tensor: + # Compute the mean across the batch dimension + mean_act = torch.mean(activations, dim=0) + + # Subtract mean from activations and square the result + deviations = (activations - mean_act) ** 2 + + # Outer product along feature dimension for each batch + # This computes the pairwise squared differences efficiently + loss = torch.einsum('bi,bj->bij', deviations, deviations) + + # Sum over all batches and features, exclude the diagonal elements + # Diagonal elements correspond to the squared terms which we want to avoid + batch_size, n_features = activations.shape + loss = torch.sum(loss) - \ + torch.sum(torch.einsum('bii->b', loss)) / 2 + + # Normalize the loss + loss = loss / (batch_size * n_features * (n_features - 1)) + + return loss + + assert current_activations_with_grad.requires_grad == True + pos_loss = generate_loss(current_activations_with_grad) + return pos_loss + # @profile(stdout=False, filename='baseline.prof', # skip=Settings.new().model.skip_profiling) + def train_layer(self, # type: ignore[override] input_data: TrainInputData, label_data: TrainLabelData, @@ -421,44 +518,115 @@ def train_layer(self, # type: ignore[override] self.optimizer.zero_grad() pos_activations = None - neg_activations = None + # neg_activations = None if input_data is not None and label_data is not None: - (pos_input, neg_input) = input_data - (pos_labels, neg_labels) = label_data + try: + (pos_input, neg_input) = input_data + (pos_labels, neg_labels) = label_data + except ValueError: + pos_input = input_data + pos_labels = label_data + pos_activations = self.forward( ForwardMode.PositiveData, pos_input, pos_labels, should_damp) - neg_activations = self.forward( - ForwardMode.NegativeData, neg_input, neg_labels, should_damp) + # neg_activations = self.forward( + # ForwardMode.NegativeData, neg_input, neg_labels, should_damp) elif input_data is not None: - (pos_input, neg_input) = input_data + try: + (pos_input, neg_input) = input_data + except ValueError: + pos_input = input_data pos_activations = self.forward( ForwardMode.PositiveData, pos_input, None, should_damp) - neg_activations = self.forward( - ForwardMode.NegativeData, neg_input, None, should_damp) + # neg_activations = self.forward( + # ForwardMode.NegativeData, neg_input, None, should_damp) elif label_data is not None: - (pos_labels, neg_labels) = label_data + try: + (pos_labels, neg_labels) = label_data + except ValueError: + pos_labels = label_data pos_activations = self.forward( ForwardMode.PositiveData, None, pos_labels, should_damp) - neg_activations = self.forward( - ForwardMode.NegativeData, None, neg_labels, should_damp) + # neg_activations = self.forward( + # ForwardMode.NegativeData, None, neg_labels, should_damp) else: pos_activations = self.forward( ForwardMode.PositiveData, None, None, should_damp) - neg_activations = self.forward( - ForwardMode.NegativeData, None, None, should_damp) + # neg_activations = self.forward( + # ForwardMode.NegativeData, None, None, should_damp) pos_badness = layer_activations_to_badness(pos_activations) - neg_badness = layer_activations_to_badness(neg_activations) - - # Loss function equivelent to: - # plot3d log(1 + exp(-n + 1)) + log(1 + exp(p - 1)) for n=0 to 3, p=0 - # to 3 - layer_loss: Tensor = F.softplus(torch.cat([ - (-1 * neg_badness) + self.settings.model.loss_threshold, + ff_layer_loss: Tensor = F.softplus( pos_badness - self.settings.model.loss_threshold - ])).mean() + ).mean() + ff_layer_loss = self.settings.model.loss_scale_ff * ff_layer_loss + + # y = 4exp(-log(x+1.5)) + # ff_layer_loss_min: Tensor = torch.sqrt(torch.square(pos_activations) + 0.001).mean() + # ff_layer_loss_min = 4 * torch.exp(-torch.log(ff_layer_loss_min + self.settings.model.loss_threshold)) + + lpl_loss_predictive: Tensor = self.settings.model.loss_scale_predictive * \ + self.generate_lpl_loss_predictive(pos_activations) + lpl_loss_hebbian: Tensor = self.settings.model.loss_scale_hebbian * \ + self.generate_lpl_loss_hebbian(pos_activations) + lpl_loss_decorrelative: Tensor = self.settings.model.loss_scale_decorrelative * \ + self.generate_lpl_loss_decorrelative(pos_activations) + + assert ff_layer_loss.requires_grad == True + # assert ff_layer_loss_min.requires_grad == True + assert lpl_loss_predictive.requires_grad == True + assert lpl_loss_hebbian.requires_grad == True + assert lpl_loss_decorrelative.requires_grad == True + + # if random.random() < 0.005: + # # print("pos_act: ", pos_activations) + # print("ff_layer_loss: ", ff_layer_loss) + # print("lpl_loss_predictive: ", lpl_loss_predictive) + # print("lpl_loss_hebbian: ", lpl_loss_hebbian) + # print("lpl_loss_decorrelative: ", lpl_loss_decorrelative) + # print() + + layer_loss: Tensor = ff_layer_loss + lpl_loss_predictive + \ + lpl_loss_hebbian + lpl_loss_decorrelative + # layer_loss: Tensor = ff_layer_loss + ff_layer_loss_min + + # print(ff_layer_loss) + # print(lpl_loss_predictive) + # print(lpl_loss_hebbian) + # print(lpl_loss_decorrelative) + + # important block + # + # print("ff_layer_loss: ", ff_layer_loss) + # print("lpl_loss_predictive: ", lpl_loss_predictive) + # print("lpl_loss_hebbian: ", lpl_loss_hebbian) + # print("lpl_loss_decorrelative: ", lpl_loss_decorrelative) + # print(self.size) + + # print(pos_activations[0]) + # input() + + # print(layer_loss) + # for name, param in self.named_parameters(): + # print(name, param.grad) + # optimizer_params = [] + # for d in self.optimizer.param_groups: + # for tensor in d['params']: + # print(tensor.shape) + # optimizer_params.append(tensor) + # print("optimizer params") + # input() + # print() layer_loss.backward() + # params = dict(self.named_parameters()) + # dot = make_dot(layer_loss, params=params) + # dot.render('computation_graph', format='png') + # input() + + # print all grads of parameters + # for name, param in self.named_parameters(): + # print(name, param.grad) self.optimizer.step() return cast(float, layer_loss.item()) @@ -541,15 +709,14 @@ def forward(self, mode: ForwardMode, data: torch.Tensor, labels: torch.Tensor, s Activations, previous_layer.predict_activations).previous prev_act = cast(Activations, self.predict_activations).previous - prev_layer_prev_timestep_activations = prev_layer_prev_timestep_activations.detach() + prev_layer_prev_timestep_activations = prev_layer_prev_timestep_activations prev_layer_stdized = standardize_layer_activations( prev_layer_prev_timestep_activations, self.settings.model.epsilon) - next_layer_prev_timestep_activations = next_layer_prev_timestep_activations.detach() + next_layer_prev_timestep_activations = next_layer_prev_timestep_activations next_layer_stdized = standardize_layer_activations( next_layer_prev_timestep_activations, self.settings.model.epsilon) - prev_act = prev_act.detach() prev_act_stdized = standardize_layer_activations( prev_act, self.settings.model.epsilon) @@ -579,7 +746,6 @@ def forward(self, mode: ForwardMode, data: torch.Tensor, labels: torch.Tensor, s assert self.predict_activations is not None prev_act = cast(Activations, self.predict_activations).previous - prev_act = prev_act.detach() prev_act_stdized = standardize_layer_activations( prev_act, self.settings.model.epsilon) @@ -612,11 +778,10 @@ def forward(self, mode: ForwardMode, data: torch.Tensor, labels: torch.Tensor, s Activations, next_layer.predict_activations).previous prev_act = cast(Activations, self.predict_activations).previous - next_layer_prev_timestep_activations = next_layer_prev_timestep_activations.detach() + next_layer_prev_timestep_activations = next_layer_prev_timestep_activations next_layer_stdized = standardize_layer_activations( next_layer_prev_timestep_activations, self.settings.model.epsilon) - prev_act = prev_act.detach() prev_act_stdized = standardize_layer_activations( prev_act, self.settings.model.epsilon) @@ -649,11 +814,10 @@ def forward(self, mode: ForwardMode, data: torch.Tensor, labels: torch.Tensor, s Activations, previous_layer.predict_activations).previous prev_act = cast(Activations, self.predict_activations).previous - prev_layer_prev_timestep_activations = prev_layer_prev_timestep_activations.detach() + prev_layer_prev_timestep_activations = prev_layer_prev_timestep_activations prev_layer_stdized = standardize_layer_activations( prev_layer_prev_timestep_activations, self.settings.model.epsilon) - prev_act = prev_act.detach() prev_act_stdized = standardize_layer_activations( prev_act, self.settings.model.epsilon) @@ -688,12 +852,12 @@ def forward(self, mode: ForwardMode, data: torch.Tensor, labels: torch.Tensor, s if mode == ForwardMode.PositiveData: assert self.pos_activations is not None - self.pos_activations.current = new_activation + self.pos_activations.current = new_activation.detach() elif mode == ForwardMode.NegativeData: assert self.neg_activations is not None - self.neg_activations.current = new_activation + self.neg_activations.current = new_activation.detach() elif mode == ForwardMode.PredictData: assert self.predict_activations is not None - self.predict_activations.current = new_activation + self.predict_activations.current = new_activation.detach() return new_activation diff --git a/RecurrentFF/model/inner_layers.py b/RecurrentFF/model/inner_layers.py index ad7b9e0..ba90a7e 100644 --- a/RecurrentFF/model/inner_layers.py +++ b/RecurrentFF/model/inner_layers.py @@ -255,7 +255,6 @@ def advance_layers_train( def advance_layers_forward( self, - mode: ForwardMode, input_data: torch.Tensor, label_data: torch.Tensor, should_damp: bool) -> None: @@ -296,15 +295,16 @@ def advance_layers_forward( this is because the overhead of creating threads is not worth it for the small amount of computation done in each thread. """ + # TODOPRE: save and restore params for i, layer in enumerate(self.layers): if i == 0 and len(self.layers) == 1: - layer.forward(mode, input_data, label_data, should_damp) + layer.train_layer(input_data, label_data, should_damp) elif i == 0: - layer.forward(mode, input_data, None, should_damp) + layer.train_layer(input_data, None, should_damp) elif i == len(self.layers) - 1: - layer.forward(mode, None, label_data, should_damp) + layer.train_layer(None, label_data, should_damp) else: - layer.forward(mode, None, None, should_damp) + layer.train_layer(None, None, should_damp) for layer in self.layers: layer.advance_stored_activations() diff --git a/RecurrentFF/model/model.py b/RecurrentFF/model/model.py index 78243dc..f5ba22a 100644 --- a/RecurrentFF/model/model.py +++ b/RecurrentFF/model/model.py @@ -17,7 +17,6 @@ from RecurrentFF.model.inner_layers import InnerLayers, LayerMetrics from RecurrentFF.util import ( Activations, - ForwardMode, LatentAverager, TrainInputData, TrainLabelData, @@ -172,6 +171,9 @@ def train_model(self, train_loader: torch.utils.data.DataLoader, test_loader: to self.train() for batch_num, (input_data, label_data) in enumerate(train_loader): + # if batch_num == 1: + # break + input_data.move_to_device_inplace(self.settings.device.device) label_data.move_to_device_inplace(self.settings.device.device) @@ -198,7 +200,7 @@ def train_model(self, train_loader: torch.utils.data.DataLoader, test_loader: to # # TODO: Fix this hacky data loader bridge format train_accuracy = self.processor.brute_force_predict( - TrainTestBridgeFormatLoader(train_loader), 10, False) # type: ignore[arg-type] + TrainTestBridgeFormatLoader(train_loader), 1, False) # type: ignore[arg-type] test_accuracy = self.processor.brute_force_predict( test_loader, 1, True) @@ -216,6 +218,8 @@ def train_model(self, train_loader: torch.utils.data.DataLoader, test_loader: to self.inner_layers.step_learning_rates() + return train_accuracy + def __train_batch( self, batch_num: int, @@ -224,22 +228,22 @@ def __train_batch( total_batch_count: int) -> Tuple[LayerMetrics, List[float], List[float]]: logging.info("Batch: " + str(batch_num)) - self.inner_layers.reset_activations(True) + # self.inner_layers.reset_activations(True) for preinit_step in range(0, self.settings.model.prelabel_timesteps): logging.debug("Preinitialization step: " + str(preinit_step)) pos_input = input_data.pos_input[0] - neg_input = input_data.neg_input[0] + # neg_input = input_data.neg_input[0] preinit_upper_clamped_tensor = self.processor.get_preinit_upper_clamped_tensor( label_data.pos_labels[0].shape) self.inner_layers.advance_layers_forward( - ForwardMode.PositiveData, pos_input, preinit_upper_clamped_tensor, False) - self.inner_layers.advance_layers_forward( - ForwardMode.NegativeData, neg_input, preinit_upper_clamped_tensor, False) + pos_input, preinit_upper_clamped_tensor, False) + # self.inner_layers.advance_layers_forward( + # neg_input, preinit_upper_clamped_tensor, False) num_layers = len(self.settings.model.hidden_sizes) layer_metrics = LayerMetrics(num_layers) diff --git a/RecurrentFF/search/new_search.py b/RecurrentFF/search/new_search.py new file mode 100644 index 0000000..f6ef5da --- /dev/null +++ b/RecurrentFF/search/new_search.py @@ -0,0 +1,188 @@ +import logging +import random +import time +from typing import TextIO +import torch +import wandb +from RecurrentFF.benchmarks.mnist.mnist import MNIST_loaders +from RecurrentFF.model.model import RecurrentFFNet +from RecurrentFF.settings import Settings + +from RecurrentFF.util import set_logging + + +# DATA_SIZE = 784 +# NUM_CLASSES = 10 +# TRAIN_BATCH_SIZE = 500 +# TEST_BATCH_SIZE = 500 + +DEVICE = "mps" + +NUM_SEEDS_BENCH = 1 + +datetime_str = time.strftime("%Y%m%d-%H%M%S") +RUNNING_LOG_FILENAME = f"running_log_{datetime_str}.log" + + +def objective() -> None: + wandb.init( + project="Recurrent-FF", + config={ + "architecture": "Recurrent-FF", + "dataset": "MNIST", + }, + # allow_val_change=True # TODOPRE: review this as it is silencing a warning + ) + + settings = Settings.new() + + settings.model.hidden_sizes = [100, 100, 100, 100, 100] + + settings.model.ff_rmsprop.learning_rate = wandb.config.learning_rate + settings.model.ff_rmsprop.momentum = wandb.config.momentum + + settings.model.prelabel_timesteps = wandb.config.prelabel_timesteps + settings.data_config.iterations = wandb.config.iterations + + settings.model.loss_scale_ff = wandb.config.loss_scale_ff + settings.model.loss_scale_predictive = wandb.config.loss_scale_predictive + settings.model.loss_scale_hebbian = wandb.config.loss_scale_hebbian + settings.model.loss_scale_decorrelative = wandb.config.loss_scale_decorrelative + + settings.model.damping_factor = wandb.config.damping_factor + + # layer_sizes = wandb.config.layer_sizes + # learning_rate = wandb.config.learning_rate + # dt = wandb.config.dt + # exc_to_inhib_conn_c = wandb.config.exc_to_inhib_conn_c + # exc_to_inhib_conn_sigma_squared = wandb.config.exc_to_inhib_conn_sigma_squared + # percentage_inhibitory = wandb.config.percentage_inhibitory + # decay_beta = wandb.config.decay_beta + # tau_mean = wandb.config.tau_mean + # tau_var = wandb.config.tau_var + # tau_stdp = wandb.config.tau_stdp + # tau_rise_alpha = wandb.config.tau_rise_alpha + # tau_fall_alpha = wandb.config.tau_fall_alpha + # tau_rise_epsilon = wandb.config.tau_rise_epsilon + # tau_fall_epsilon = wandb.config.tau_fall_epsilon + # threshold_scale = wandb.config.threshold_scale + # threshold_decay = wandb.config.threshold_decay + + # sum layer sizes to get total neurons + # total_neurons = sum(layer_sizes) + # layer_sparsity = NUM_NEURONS_CONNECT_ACROSS_LAYERS / total_neurons + + run_settings = f""" + running with: + settings: {settings.model_dump()} + """ + # run_settings = f""" + # running with: + # layer_sizes: {layer_sizes} + # learning_rate: {learning_rate} + # dt: {dt} + # percentage_inhibitory: {percentage_inhibitory} + # exc_to_inhib_conn_c: {exc_to_inhib_conn_c} + # exc_to_inhib_conn_sigma_squared: {exc_to_inhib_conn_sigma_squared} + # layer_sparsity: {layer_sparsity} + # decay_beta: {decay_beta}, + # tau_mean: {tau_mean}, + # tau_var: {tau_var}, + # tau_stdp: {tau_stdp}, + # tau_rise_alpha: {tau_rise_alpha}, + # tau_fall_alpha: {tau_fall_alpha}, + # tau_rise_epsilon: {tau_rise_epsilon}, + # tau_fall_epsilon: {tau_fall_epsilon}, + # threshold_scale: {threshold_scale}, + # threshold_decay: {threshold_decay}, + # """ + # logging.info(run_settings) + + with open(RUNNING_LOG_FILENAME, "a") as running_log: + running_log.write(f"{run_settings}") + running_log.flush() + + cum_pass_rate = 0 + for _ in range(NUM_SEEDS_BENCH): + pass_rate = bench_specific_seed( + running_log, + settings + ) + wandb.log({"train_accuracy": pass_rate}) + cum_pass_rate += pass_rate + + running_log.write( + run_settings + f"train_accuracy: {cum_pass_rate / NUM_SEEDS_BENCH}\n\n======================================\ + =========================================") + running_log.flush() + + wandb.log({"average_image_predict_success": cum_pass_rate / NUM_SEEDS_BENCH}) + + +# @profile(stdout=False, filename='baseline.prof', +# skip=True) +def bench_specific_seed(running_log: TextIO, + settings: Settings + ) -> float: + rand = random.randint(1000, 9999) + torch.manual_seed(rand) + running_log.write(f"Seed: {rand}\n") + + # settings = Settings( + # layer_sizes=layer_sizes, + # data_size=dataset.num_classes, + # batch_size=BATCH_SIZE, + # learning_rate=learning_rate, + # epochs=10, + # encode_spike_trains=ENCODE_SPIKE_TRAINS, + # dt=dt, + # percentage_inhibitory=percentage_inhibitory, + # exc_to_inhib_conn_c=exc_to_inhib_conn_c, + # exc_to_inhib_conn_sigma_squared=exc_to_inhib_conn_sigma_squared, + # layer_sparsity=layer_sparsity, + # decay_beta=decay_beta, + # tau_mean=tau_mean, + # tau_var=tau_var, + # tau_stdp=tau_stdp, + # tau_rise_alpha=tau_rise_alpha, + # tau_fall_alpha=tau_fall_alpha, + # tau_rise_epsilon=tau_rise_epsilon, + # tau_fall_epsilon=tau_fall_epsilon, + # threshold_scale=threshold_scale, + # threshold_decay=threshold_decay, + # device=torch.device(DEVICE) + # ) + + # train_dataloader = DataLoader(dataset, batch_size=settings.batch_size, shuffle=False) + + net = RecurrentFFNet(settings).to(settings.device.device) + + train_loader, test_loader = MNIST_loaders( + settings.data_config.train_batch_size, settings.data_config.test_batch_size) + + train_accuracy = net.train_model(train_loader, test_loader) + + message = f"""--------------------------------- + train_accuracy: {train_accuracy} + --------------------------------- + """ + running_log.write(message) + running_log.flush() + logging.info(message) + + return train_accuracy + + +if __name__ == "__main__": + torch.autograd.set_detect_anomaly(True) + torch.set_printoptions(precision=10, sci_mode=False) + set_logging() + + running_log = open(RUNNING_LOG_FILENAME, "w") + message = f"Sweep logs. Current datetime: {time.ctime()}\n" + running_log.write(message) + running_log.close() + logging.debug(message) + + sweep_id = "and-rewsmith/Recurrent-FF/04e3a5rr" + wandb.agent(sweep_id, function=objective) diff --git a/RecurrentFF/search/sweep_config.yaml b/RecurrentFF/search/sweep_config.yaml new file mode 100644 index 0000000..72433d8 --- /dev/null +++ b/RecurrentFF/search/sweep_config.yaml @@ -0,0 +1,33 @@ +program: new_search.py +method: bayes +metric: + goal: maximize + name: train_accuracy +parameters: + learning_rate: + min: 0.0000005 + max: 0.0005 + momentum: + min: 0.0 + max: 1.0 + prelabel_timesteps: + min: 3 + max: 10 + iterations: + min: 10 + max: 20 + loss_scale_ff: + min: 0.0 + max: 1.0 + loss_scale_predictive: + min: 0.0 + max: 1.0 + loss_scale_hebbian: + min: 0.0 + max: 1.0 + loss_scale_decorrelative: + min: 0.0 + max: 1.0 + damping_factor: + min: 0.1 + max: 0.9 \ No newline at end of file diff --git a/RecurrentFF/settings.py b/RecurrentFF/settings.py index b97a233..b71f5f5 100644 --- a/RecurrentFF/settings.py +++ b/RecurrentFF/settings.py @@ -55,6 +55,10 @@ class Model(BaseModel): should_replace_neg_data: bool should_load_weights: bool dropout: float + loss_scale_ff: float + loss_scale_predictive: float + loss_scale_hebbian: float + loss_scale_decorrelative: float lr_step_size: int lr_gamma: float diff --git a/RecurrentFF/util.py b/RecurrentFF/util.py index 1fcaca1..e2aff56 100644 --- a/RecurrentFF/util.py +++ b/RecurrentFF/util.py @@ -77,7 +77,7 @@ def __iter__(self) -> Generator[torch.Tensor, None, None]: yield self.previous def advance(self) -> None: - self.previous = self.current + self.previous = self.current.clone() class ForwardMode(Enum): diff --git a/computation_graph b/computation_graph new file mode 100644 index 0000000..ec9a32a --- /dev/null +++ b/computation_graph @@ -0,0 +1,143 @@ +digraph { + graph [size="19.8,19.8"] + node [align=left fontname=monospace fontsize=10 height=0.2 ranksep=0.1 shape=box style=filled] + 5028592144 [label=" + ()" fillcolor=darkolivegreen1] + 5028444480 [label=AddBackward0] + 5028444624 -> 5028444480 + 5028444624 [label=AddBackward0] + 5028444528 -> 5028444624 + 5028444528 [label=AddBackward0] + 5028444912 -> 5028444528 + 5028444912 [label=MulBackward0] + 5028445056 -> 5028444912 + 5028445056 [label=MeanBackward0] + 5028445152 -> 5028445056 + 5028445152 [label=SoftplusBackward0] + 5028445248 -> 5028445152 + 5028445248 [label=SubBackward0] + 5028445344 -> 5028445248 + 5028445344 [label=MeanBackward1] + 5028445440 -> 5028445344 + 5028445440 [label=PowBackward0] + 5028445488 -> 5028445440 + 5028445488 [label=LeakyReluBackward0] + 5028445632 -> 5028445488 + 5028445632 [label=AddBackward0] + 5028445776 -> 5028445632 + 5028445776 [label=AddBackward0] + 5028446016 -> 5028445776 + 5028446016 [label=LinearBackward0] + 5028446160 -> 5028446016 + 5028007568 [label="forward_linear.weight + (701, 784)" fillcolor=lightblue] + 5028007568 -> 5028446160 + 5028446160 [label=AccumulateGrad] + 5028446112 -> 5028446016 + 5028005648 [label="forward_linear.bias + (701)" fillcolor=lightblue] + 5028005648 -> 5028446112 + 5028446112 [label=AccumulateGrad] + 5028445968 -> 5028445776 + 5028445968 [label=MulBackward0] + 5028446064 -> 5028445968 + 5028446064 [label=LinearBackward0] + 5028856112 -> 5028446064 + 5028009104 [label="backward_linear.weight + (701, 702)" fillcolor=lightblue] + 5028009104 -> 5028856112 + 5028856112 [label=AccumulateGrad] + 5028856064 -> 5028446064 + 5028009296 [label="backward_linear.bias + (701)" fillcolor=lightblue] + 5028009296 -> 5028856064 + 5028856064 [label=AccumulateGrad] + 5028445728 -> 5028445632 + 5028445728 [label=LinearBackward0] + 5028445920 -> 5028445728 + 5028008432 [label="lateral_linear.weight + (701, 701)" fillcolor=lightblue] + 5028008432 -> 5028445920 + 5028445920 [label=AccumulateGrad] + 5028855872 -> 5028445728 + 5028008816 [label="lateral_linear.bias + (701)" fillcolor=lightblue] + 5028008816 -> 5028855872 + 5028855872 [label=AccumulateGrad] + 5028444864 -> 5028444528 + 5028444864 [label=MulBackward0] + 5028445200 -> 5028444864 + 5028445200 [label=DivBackward0] + 5028445392 -> 5028445200 + 5028445392 [label=SumBackward1] + 5028445584 -> 5028445392 + 5028445584 [label=SumBackward1] + 5028445872 -> 5028445584 + 5028445872 [label=PowBackward0] + 5028856208 -> 5028445872 + 5028856208 [label=SubBackward0] + 5028445488 -> 5028856208 + 5028444768 -> 5028444624 + 5028444768 [label=MulBackward0] + 5028445296 -> 5028444768 + 5028445296 [label=DivBackward0] + 5028445008 -> 5028445296 + 5028445008 [label=NegBackward0] + 5028444816 -> 5028445008 + 5028444816 [label=SumBackward0] + 5028856016 -> 5028444816 + 5028856016 [label=LogBackward0] + 5028856400 -> 5028856016 + 5028856400 [label=DivBackward0] + 5028856496 -> 5028856400 + 5028856496 [label=SumBackward1] + 5028856592 -> 5028856496 + 5028856592 [label=PowBackward0] + 5028856688 -> 5028856592 + 5028856688 [label=SubBackward0] + 5028445488 -> 5028856688 + 5028856784 -> 5028856688 + 5028856784 [label=MeanBackward1] + 5028445488 -> 5028856784 + 5028444672 -> 5028444480 + 5028444672 [label=MulBackward0] + 5028444960 -> 5028444672 + 5028444960 [label=DivBackward0] + 5028444384 -> 5028444960 + 5028444384 [label=SubBackward0] + 5028856448 -> 5028444384 + 5028856448 [label=SumBackward0] + 5028856736 -> 5028856448 + 5028856736 [label=MulBackward0] + 5028856832 -> 5028856736 + 5028856832 [label=PermuteBackward0] + 5028856976 -> 5028856832 + 5028856976 [label=UnsqueezeBackward0] + 5028857072 -> 5028856976 + 5028857072 [label=PowBackward0] + 5028857168 -> 5028857072 + 5028857168 [label=SubBackward0] + 5028445488 -> 5028857168 + 5028857264 -> 5028857168 + 5028857264 [label=MeanBackward1] + 5028445488 -> 5028857264 + 5028856880 -> 5028856736 + 5028856880 [label=PermuteBackward0] + 5028857120 -> 5028856880 + 5028857120 [label=UnsqueezeBackward0] + 5028857072 -> 5028857120 + 5028856352 -> 5028444384 + 5028856352 [label=DivBackward0] + 5028857024 -> 5028856352 + 5028857024 [label=SumBackward0] + 5028857216 -> 5028857024 + 5028857216 [label=SumBackward1] + 5028857312 -> 5028857216 + 5028857312 [label=PermuteBackward0] + 5028857408 -> 5028857312 + 5028857408 [label=PermuteBackward0] + 5028857504 -> 5028857408 + 5028857504 [label=DiagonalBackward0] + 5028856736 -> 5028857504 + 5028444480 -> 5028592144 +} diff --git a/computation_graph.png b/computation_graph.png new file mode 100644 index 0000000..dec2085 Binary files /dev/null and b/computation_graph.png differ diff --git a/config.toml b/config.toml index ac409bb..620bfad 100644 --- a/config.toml +++ b/config.toml @@ -2,8 +2,8 @@ device = "mps" [model] -hidden_sizes = [700, 700, 700, 700, 700] -epochs = 30000 +hidden_sizes = [100, 100, 100, 100, 100] +epochs = 25 prelabel_timesteps = 10 loss_threshold = 1.5 damping_factor = 0.7 @@ -18,6 +18,10 @@ should_load_weights = false lr_step_size = 10000 lr_gamma = 0.90 dropout = 0.0 +loss_scale_ff = 1 +loss_scale_predictive = 0 +loss_scale_hebbian = 0 +loss_scale_decorrelative = 0 [model.ff_rmsprop] momentum = 0.0 @@ -38,3 +42,11 @@ learning_rate = 0.0001 [model.classifier_adadelta] learning_rate = 0.00001 + +[data_config] +data_size = 784 +num_classes = 10 +train_batch_size = 500 +test_batch_size = 500 +dataset = "MNIST" +iterations = 10