| 283 | |
| 284 | class KNNSaverMulti(object): |
| 285 | def __init__(self, dstore_dir, dimension, knn_keytype=None, knn_gpu=False, training_args=None, eval_subset=None, accelerator=None): |
| 286 | self.eval_subset = eval_subset |
| 287 | self.dstore_dir = dstore_dir |
| 288 | self.dimension = dimension |
| 289 | self.knn_keytype = KEY_TYPE.last_ffn_input if knn_keytype is None else knn_keytype |
| 290 | self.training_args = training_args |
| 291 | |
| 292 | # Multi-GPU settings |
| 293 | self.world_size = training_args.world_size if training_args else 1 |
| 294 | self.process_index = training_args.local_process_index if training_args else 0 |
| 295 | self.device = training_args.device if training_args else torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| 296 | self.accelerator = accelerator |
| 297 | |
| 298 | self.model = None |
| 299 | self.activation_capturer = None |
| 300 | self.is_encoder_decoder = None |
| 301 | self.dstore_idx = 0 |
| 302 | self.hook_handles = [] |
| 303 | self.knn_gpu = knn_gpu |
| 304 | |
| 305 | if self.process_index ==0: |
| 306 | if not os.path.exists(self.dstore_dir): |
| 307 | logger.info(f"Creating directory {self.dstore_dir} for storing the datastore.") |
| 308 | os.makedirs(self.dstore_dir, exist_ok=True) |
| 309 | |
| 310 | def _get_arrow_file_path(self): |
| 311 | """Get the Arrow file path (single file for all processes)""" |