forked from HoagyC/sparse_coding
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbasic_l1_sweep.py
More file actions
135 lines (100 loc) · 5.01 KB
/
Copy pathbasic_l1_sweep.py
File metadata and controls
135 lines (100 loc) · 5.01 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
import torch
import torchopt
import numpy as np
from big_sweep import ensemble_train_loop, unstacked_to_learned_dicts
from autoencoders.sae_ensemble import FunctionalTiedSAE
from autoencoders.ensemble import FunctionalEnsemble
import os
import tqdm
from utils import dotdict
class ProgressBarCursed:
def __init__(self, total, chunk_idx, n_chunks, epoch_idx, n_repetitions):
if n_repetitions > 1:
desc = "Epoch {}/{} - Chunk {}/{}".format(epoch_idx+1, n_repetitions, chunk_idx+1, n_chunks)
else:
desc = "Chunk {}/{}".format(chunk_idx+1, n_chunks)
self.bar = tqdm.tqdm(total=total, desc=desc)
self._value = 0
@property
def value(self):
return self._value
@value.setter
def value(self, v):
self.bar.update(v - self._value)
self._value = v
def basic_l1_sweep(
dataset_dir, output_dir,
ratio, l1_values=np.logspace(-4, -2, 16), batch_size=256,
device="cuda", adam_kwargs={"lr": 1e-3},
n_repetitions=1,
save_after_every=False,
):
# get dataset size
# check that dataset_dir/0.pt exists
assert os.path.exists(os.path.join(dataset_dir, '0.pt')), "Dataset not found at {}".format(dataset_dir)
dataset = torch.load(os.path.join(dataset_dir, '0.pt'))
activation_dim = dataset.shape[1]
latent_dim = int(activation_dim * ratio)
del dataset
# create models
print(f"Initializing {len(l1_values)} models with latent dimension {latent_dim}...")
models = [FunctionalTiedSAE.init(activation_dim, latent_dim, l1, device=device) for l1 in l1_values]
ensemble = FunctionalEnsemble(
models, FunctionalTiedSAE,
torchopt.adam, adam_kwargs,
device=device
)
args = {
"batch_size": batch_size,
"device": device,
"dict_size": latent_dim,
}
print("Training...")
n_chunks = len(os.listdir(dataset_dir))
os.makedirs(output_dir, exist_ok=True)
for epoch_idx in range(n_repetitions):
chunk_order = np.random.permutation(n_chunks)
for chunk_idx, chunk in enumerate(chunk_order):
assert os.path.exists(os.path.join(dataset_dir, '{}.pt'.format(chunk))), "Chunk not found at {}".format(os.path.join(dataset_dir, '{}.pt'.format(chunk)))
dataset = torch.load(os.path.join(dataset_dir, '{}.pt'.format(chunk))).to(dtype=torch.float32)
dataset.pin_memory()
sampler = torch.utils.data.BatchSampler(
torch.utils.data.RandomSampler(range(dataset.shape[0])),
batch_size=batch_size,
drop_last=False,
)
bar = ProgressBarCursed(len(sampler), chunk_idx, n_chunks, epoch_idx, n_repetitions)
cfg = dotdict({
"use_wandb": False,
})
ensemble_train_loop(ensemble, cfg, args, "ensemble", sampler, dataset, bar)
if save_after_every:
learned_dicts = unstacked_to_learned_dicts(ensemble, args, ["dict_size"], ["l1_alpha"])
torch.save(learned_dicts, os.path.join(output_dir, f"learned_dicts_epoch_{epoch_idx}_chunk_{chunk_idx}.pt"))
if not save_after_every:
learned_dicts = unstacked_to_learned_dicts(ensemble, args, ["dict_size"], ["l1_alpha"])
torch.save(learned_dicts, os.path.join(output_dir, f"learned_dicts_epoch_{epoch_idx}.pt"))
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Train an ensemble of SAEs with different L1 penalties.")
parser.add_argument("--dataset_dir", type=str, help="Path to directory containing dataset chunks.")
parser.add_argument("--output_dir", type=str, help="Path to directory to save learned dictionaries.")
parser.add_argument("--ratio", type=float, help="Ratio of latent to activation dimension.")
parser.add_argument("--l1_value_min", type=float, default=-4, help="Minimum L1 penalty value (log base 10).")
parser.add_argument("--l1_value_max", type=float, default=-2, help="Maximum L1 penalty value (log base 10).")
parser.add_argument("--l1_value_n", type=int, default=16, help="Number of L1 penalty values to try.")
parser.add_argument("--batch_size", type=int, default=256, help="Batch size.")
parser.add_argument("--device", type=str, default="cuda", help="Device to use.")
parser.add_argument("--adam_lr", type=float, default=1e-3, help="Adam learning rate.")
parser.add_argument("--n_repetitions", type=int, default=1, help="Number of epochs to train for.")
parser.add_argument("--save_after_every", action="store_true", help="Save learned dictionaries after every chunk instead of every epoch.")
args = parser.parse_args()
#l1_values = list(np.logspace(args.l1_value_min, args.l1_value_max, args.l1_value_n))
l1_values = [0, 1e-3, 3e-4, 1e-4]
basic_l1_sweep(
args.dataset_dir, args.output_dir,
args.ratio, l1_values, args.batch_size,
args.device, {"lr": args.adam_lr},
args.n_repetitions,
args.save_after_every
)