|
AQIT 0.1.0
|
Functions | |
| chunk_paths (int layer) | |
| Path | norm_cache_path (int layer) |
| Path | sae_save_path (int layer) |
| collect_activations (model, int layer) | |
| compute_norm (int layer) | |
| train_layer (int layer) | |
Variables | |
| str | MODEL_NAME = "meta-llama/Llama-3.2-1B-Instruct" |
| int | N_LAYERS = 16 |
| int | N_FEATURES = 32768 |
| int | D_MODEL = 2048 |
| float | L1_COEFF = 10.0 |
| int | LR = 1e-4 |
| int | BATCH_SIZE = 4096 |
| int | N_TOKENS = 2_000_000 |
| int | SEQ_LEN = 64 |
| int | COLLECT_BATCH = 32 |
| int | CHUNK_SIZE = 100_000 |
| int | SAVE_EVERY = 2000 |
| str | SAE_DIR = Path(__file__).parent / "sae" |
| str | ACTS_DIR = Path(__file__).parent / "sae" / "acts" |
| DEVICE = resolve_compute_device() | |
| parser = argparse.ArgumentParser() | |
| type | |
| str | |
| default | |
| help | |
| args = parser.parse_args() | |
| str | layers = "all" else [int(x) for x in args.layers.split(",")] |
| flush | |
| model = HookedTransformer.from_pretrained(MODEL_NAME) | |
| chunk_paths | ( | int | layer | ) |
Definition at line 41 of file train_sae.py.
Referenced by collect_activations(), compute_norm(), and train_layer().
| collect_activations | ( | model, | |
| int | layer ) |
Definition at line 53 of file train_sae.py.
References chunk_paths().
| compute_norm | ( | int | layer | ) |
Definition at line 105 of file train_sae.py.
References chunk_paths(), and norm_cache_path().
Referenced by train_layer().
| Path norm_cache_path | ( | int | layer | ) |
Definition at line 45 of file train_sae.py.
Referenced by compute_norm().
| Path sae_save_path | ( | int | layer | ) |
Definition at line 49 of file train_sae.py.
Referenced by train_layer().
| train_layer | ( | int | layer | ) |
Definition at line 131 of file train_sae.py.
References chunk_paths(), compute_norm(), and sae_save_path().
| str aquin.compute.train_sae.ACTS_DIR = Path(__file__).parent / "sae" / "acts" |
Definition at line 37 of file train_sae.py.
| aquin.compute.train_sae.args = parser.parse_args() |
Definition at line 196 of file train_sae.py.
| int aquin.compute.train_sae.BATCH_SIZE = 4096 |
Definition at line 30 of file train_sae.py.
| int aquin.compute.train_sae.CHUNK_SIZE = 100_000 |
Definition at line 34 of file train_sae.py.
| int aquin.compute.train_sae.COLLECT_BATCH = 32 |
Definition at line 33 of file train_sae.py.
| int aquin.compute.train_sae.D_MODEL = 2048 |
Definition at line 27 of file train_sae.py.
| aquin.compute.train_sae.default |
Definition at line 194 of file train_sae.py.
| aquin.compute.train_sae.DEVICE = resolve_compute_device() |
Definition at line 38 of file train_sae.py.
| aquin.compute.train_sae.flush |
Definition at line 200 of file train_sae.py.
| aquin.compute.train_sae.help |
Definition at line 195 of file train_sae.py.
| float aquin.compute.train_sae.L1_COEFF = 10.0 |
Definition at line 28 of file train_sae.py.
| str aquin.compute.train_sae.layers = "all" else [int(x) for x in args.layers.split(",")] |
Definition at line 198 of file train_sae.py.
| int aquin.compute.train_sae.LR = 1e-4 |
Definition at line 29 of file train_sae.py.
| aquin.compute.train_sae.model = HookedTransformer.from_pretrained(MODEL_NAME) |
Definition at line 201 of file train_sae.py.
| str aquin.compute.train_sae.MODEL_NAME = "meta-llama/Llama-3.2-1B-Instruct" |
Definition at line 24 of file train_sae.py.
| int aquin.compute.train_sae.N_FEATURES = 32768 |
Definition at line 26 of file train_sae.py.
| int aquin.compute.train_sae.N_LAYERS = 16 |
Definition at line 25 of file train_sae.py.
| int aquin.compute.train_sae.N_TOKENS = 2_000_000 |
Definition at line 31 of file train_sae.py.
| aquin.compute.train_sae.parser = argparse.ArgumentParser() |
Definition at line 193 of file train_sae.py.
| str aquin.compute.train_sae.SAE_DIR = Path(__file__).parent / "sae" |
Definition at line 36 of file train_sae.py.
| int aquin.compute.train_sae.SAVE_EVERY = 2000 |
Definition at line 35 of file train_sae.py.
| int aquin.compute.train_sae.SEQ_LEN = 64 |
Definition at line 32 of file train_sae.py.
| aquin.compute.train_sae.str |
Definition at line 194 of file train_sae.py.
| aquin.compute.train_sae.type |
Definition at line 194 of file train_sae.py.