Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -16,4 +16,5 @@ sequences
cifar-data/
.DS_Store
*.parquet
.mypy_cache/
.mypy_cache/
*.log
2 changes: 1 addition & 1 deletion RecurrentFF/benchmarks/mnist/mnist.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
186 changes: 100 additions & 86 deletions RecurrentFF/model/data_scenario/static_single_class.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import copy
import logging
from typing import List, Optional, cast
from pyparsing import Iterator
Expand Down Expand Up @@ -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(
Expand Down
Loading