| 273 | logger.info(f"Finished writing to {self.output_path}") |
| 274 | |
| 275 | def parse_args(): |
| 276 | import argparse |
| 277 | parser = argparse.ArgumentParser(description='Multi-GPU kNN Search and Distribution Processing') |
| 278 | parser.add_argument('--dstore_path', type=str, required=True, |
| 279 | help='Path to input Arrow file with keys and vals') |
| 280 | parser.add_argument('--val_path', type=str, required=False, default=None, |
| 281 | help='Path to input Arrow file with vals') |
| 282 | parser.add_argument('--index_path', type=str, required=True, |
| 283 | help='Path to FAISS index file') |
| 284 | parser.add_argument('--output_path', type=str, required=True, |
| 285 | help='Path to output Arrow file') |
| 286 | parser.add_argument('--model_path', type=str, required=True, |
| 287 | help='Path to model for tokenizer loading') |
| 288 | parser.add_argument('--k', type=int, default=1024, |
| 289 | help='Number of nearest neighbors to search') |
| 290 | parser.add_argument('--knn_temp', type=float, default=1.0, |
| 291 | help='Temperature for kNN probability computation') |
| 292 | parser.add_argument('--threshold', type=float, default=0, |
| 293 | help='Temperature for kNN probability computation') |
| 294 | parser.add_argument('--probe', type=int, default=32, |
| 295 | help='Number of probes for FAISS index') |
| 296 | parser.add_argument('--batch_size', type=int, default=32, |
| 297 | help='Batch size for processing') |
| 298 | parser.add_argument('--knn_gpu', action='store_true', |
| 299 | help='Use GPU for FAISS search') |
| 300 | parser.add_argument('--ignore_first', type=bool, default=False, |
| 301 | help='whether to ignore the nearest neighbor') |
| 302 | |
| 303 | args = parser.parse_args() |
| 304 | return args |
| 305 | |
| 306 | def main(): |
| 307 | args = parse_args() |