|
122 | 122 | "\n", |
123 | 123 | "# LIT integration imports\n", |
124 | 124 | "from transformer_lens.lit import (\n", |
125 | | - " HookedTransformerLIT,\n", |
126 | | - " HookedTransformerLITConfig,\n", |
| 125 | + " TransformerLensLIT,\n", |
| 126 | + " TransformerLensLITConfig,\n", |
127 | 127 | " SimpleTextDataset,\n", |
128 | 128 | " PromptCompletionDataset,\n", |
129 | 129 | " IOIDataset,\n", |
|
155 | 155 | }, |
156 | 156 | { |
157 | 157 | "cell_type": "code", |
158 | | - "execution_count": 3, |
| 158 | + "execution_count": null, |
159 | 159 | "id": "18cbabbd", |
160 | 160 | "metadata": {}, |
161 | | - "outputs": [ |
162 | | - { |
163 | | - "name": "stdout", |
164 | | - "output_type": "stream", |
165 | | - "text": [ |
166 | | - "Loading gpt2-small...\n" |
167 | | - ] |
168 | | - }, |
169 | | - { |
170 | | - "name": "stderr", |
171 | | - "output_type": "stream", |
172 | | - "text": [ |
173 | | - "`torch_dtype` is deprecated! Use `dtype` instead!\n" |
174 | | - ] |
175 | | - }, |
176 | | - { |
177 | | - "name": "stdout", |
178 | | - "output_type": "stream", |
179 | | - "text": [ |
180 | | - "Loaded pretrained model gpt2-small into HookedTransformer\n", |
181 | | - "Loaded model: gpt2\n", |
182 | | - " Layers: 12\n", |
183 | | - " Heads: 12\n", |
184 | | - " d_model: 768\n" |
185 | | - ] |
186 | | - } |
187 | | - ], |
| 161 | + "outputs": [], |
188 | 162 | "source": [ |
189 | 163 | "# Load GPT-2 (124M parameters)\n", |
190 | 164 | "# Other options: \"gpt2-medium\", \"gpt2-large\", \"gpt2-xl\", \"EleutherAI/pythia-70m\", etc.\n", |
|
231 | 205 | ], |
232 | 206 | "source": [ |
233 | 207 | "# Configure the wrapper\n", |
234 | | - "config = HookedTransformerLITConfig(\n", |
| 208 | + "config = TransformerLensLITConfig(\n", |
235 | 209 | " max_seq_length=256, # Maximum input length\n", |
236 | 210 | " batch_size=4, # Batch size for inference\n", |
237 | 211 | " top_k=10, # Number of top predictions to show\n", |
|
243 | 217 | ")\n", |
244 | 218 | "\n", |
245 | 219 | "# Create the wrapper\n", |
246 | | - "lit_model = HookedTransformerLIT(model, config=config)\n", |
| 220 | + "lit_model = TransformerLensLIT(model, config=config)\n", |
247 | 221 | "\n", |
248 | 222 | "print(f\"Created LIT wrapper: {lit_model.description()}\")\n", |
249 | 223 | "print(f\"\\nInput spec keys: {list(lit_model.input_spec().keys())}\")\n", |
|
0 commit comments