1+ from pathlib import Path
12from utils import BASE_DIR # Need this import to set the huggingface cache directory
23import os
34import numpy as np
@@ -24,7 +25,6 @@ def __str__(self):
2425 def __repr__ (self ):
2526 return str (self )
2627
27-
2828def generate_and_save_acts (
2929 task ,
3030 names_filter ,
@@ -39,10 +39,10 @@ def generate_and_save_acts(
3939 forward_batch_size = 2
4040 num_tokens_to_generate = task .num_tokens_in_answer
4141 all_problems = task .generate_problems ()
42- output_file = task .prefix + "results.csv"
42+ output_file = task .prefix / "results.csv"
4343
4444 if save_results_csv :
45- os . makedirs ( task .prefix , exist_ok = True )
45+ task .prefix . mkdir ( parents = True , exist_ok = True )
4646 model_best_addition = "" if not save_best_logit else ", best_token"
4747 with open (output_file , "w" ) as f :
4848 f .write (
@@ -98,7 +98,7 @@ def generate_and_save_acts(
9898 print (tensors .shape )
9999 torch .save (
100100 tensors ,
101- f" { task .prefix } { save_file_prefix } { current_problem_index } .pt" ,
101+ task .prefix / f" { save_file_prefix } { current_problem_index } .pt" ,
102102 )
103103
104104 if save_results_csv :
@@ -146,7 +146,7 @@ def get_all_acts(
146146 all_problems = task .generate_problems ()
147147 all_problems_already_generated = True
148148 for i in range (len (all_problems )):
149- if not os . path . exists ( f" { task .prefix } { save_file_prefix } { i } .pt" ):
149+ if not ( task .prefix / f" { save_file_prefix } { i } .pt"). exists ( ):
150150 all_problems_already_generated = False
151151 break
152152 if not all_problems_already_generated or force_regenerate :
@@ -163,7 +163,7 @@ def get_all_acts(
163163 all_acts = []
164164 for i in range (0 , len (all_problems )):
165165 tensors = torch .load (
166- f" { task .prefix } { save_file_prefix } { i } .pt" , map_location = "cpu"
166+ task .prefix / f" { save_file_prefix } { i } .pt" , map_location = "cpu"
167167 )
168168 all_acts .append (tensors )
169169 if len (all_acts ) > 1 :
@@ -186,17 +186,17 @@ def get_acts(
186186 if save_file_prefix != "" and save_file_prefix [- 1 ] != "_" :
187187 save_file_prefix += "_"
188188 file_name = (
189- f" { task .prefix } { save_file_prefix } layer{ layer_fetch } _token{ token_fetch } .pt"
189+ task .prefix / f" { save_file_prefix } layer{ layer_fetch } _token{ token_fetch } .pt"
190190 )
191- if not os . path . exists (file_name ) or force_regenerate :
191+ if not file_name . exists () or force_regenerate :
192192 print (file_name , "not exists" )
193193 all_acts = get_all_acts (
194194 task , names_filter = names_filter , save_file_prefix = save_file_prefix
195195 )
196196 for layer in range (all_acts .shape [1 ]):
197197 for token in range (all_acts .shape [2 ]):
198198 file_name = (
199- f" { task .prefix } { save_file_prefix } layer{ layer } _token{ token } .pt"
199+ task .prefix / f" { save_file_prefix } layer{ layer } _token{ token } .pt"
200200 )
201201 torch .save (
202202 all_acts [:, layer , token , :].detach ().cpu ().clone (), file_name
@@ -218,11 +218,11 @@ def get_acts_pca(
218218 names_filter = lambda x : "resid_post" in x or "hook_embed" in x ,
219219 save_file_prefix = "" ,
220220):
221- act_file_name = f" { task .prefix } pca/ { save_file_prefix } / layer{ layer } _token{ token } _pca{ pca_k } { '_normalize' if normalize_rms else '' } .pt"
222- pca_pkl_file_name = f" { task .prefix } pca/ { save_file_prefix } / layer{ layer } _token{ token } _pca{ pca_k } { '_normalize' if normalize_rms else '' } .pkl"
223- os . makedirs ( f" { task .prefix } / pca/ { save_file_prefix } " , exist_ok = True )
221+ act_file_name = task .prefix / " pca" / save_file_prefix / f" layer{ layer } _token{ token } _pca{ pca_k } { '_normalize' if normalize_rms else '' } .pt"
222+ pca_pkl_file_name = task .prefix / " pca" / save_file_prefix / f" layer{ layer } _token{ token } _pca{ pca_k } { '_normalize' if normalize_rms else '' } .pkl"
223+ ( task .prefix / " pca" / save_file_prefix ). mkdir ( parents = True , exist_ok = True )
224224
225- if not os . path . exists (act_file_name ) or not os . path . exists (pca_pkl_file_name ):
225+ if not act_file_name . exists () or not pca_pkl_file_name . exists ():
226226 acts = get_acts (
227227 task ,
228228 layer ,
@@ -239,9 +239,9 @@ def get_acts_pca(
239239
240240
241241def get_acts_pls (task , layer , token , pls_k , normalize_rms = False ):
242- act_file_name = f" { task .prefix } / pls/ layer{ layer } _token{ token } _pls{ pls_k } { '_normalize' if normalize_rms else '' } .pt"
243- pls_pkl_file_name = f" { task .prefix } / pls/ layer{ layer } _token{ token } _pls{ pls_k } { '_normalize' if normalize_rms else '' } .pkl"
244- os . makedirs ( f" { task .prefix } / pls" , exist_ok = True )
242+ act_file_name = task .prefix / " pls" / f" layer{ layer } _token{ token } _pls{ pls_k } { '_normalize' if normalize_rms else '' } .pt"
243+ pls_pkl_file_name = task .prefix / " pls" / f" layer{ layer } _token{ token } _pls{ pls_k } { '_normalize' if normalize_rms else '' } .pkl"
244+ ( task .prefix / " pls"). mkdir ( parents = True , exist_ok = True )
245245
246246 # if not os.path.exists(act_file_name) or not os.path.exists(pls_pkl_file_name):
247247 if True :
0 commit comments