MCPcopy Create free account
hub / github.com/LUMIA-Group/MemoryDecoder / __init__

Method __init__

knn_utils/saveEmbedMulti.py:48–74  ·  view source on GitHub ↗
(self, val_file, index_file, dimension, 
            knn_sim_func=None, knn_keytype=None, knn_gpu=True,
            k=1024, lmbda=0.25, knn_temp=1.0, probe=8, local_process_index=None)

Source from the content-addressed store, hash-verified

46
47class KNNWrapperMulti(object):
48 def __init__(self, val_file, index_file, dimension,
49 knn_sim_func=None, knn_keytype=None, knn_gpu=True,
50 k=1024, lmbda=0.25, knn_temp=1.0, probe=8, local_process_index=None):
51 self.val_file = val_file
52 self.index_file = index_file
53 self.dimension = dimension
54 self.lmbda = lmbda
55 self.k = k
56 self.knn_temperature = knn_temp
57 self.probe = probe
58 self.knn_sim_func = DIST.l2
59 self.knn_keytype = KEY_TYPE.last_ffn_input
60 self.knn_gpu = knn_gpu and torch.cuda.is_available() and torch.cuda.device_count() > 0
61 self.local_process_index = local_process_index
62
63 self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
64 self.model = None
65 self.vocab_size = None
66 self.activation_capturer = None
67 self.is_encoder_decoder = None
68 self.hook_handles = []
69
70 dist_type_to_dist_func = {
71 DIST.l2: KNNWrapperMulti.l2,
72 DIST.dot: KNNWrapperMulti.dotprod,
73 }
74 self.dist_func = dist_type_to_dist_func[self.knn_sim_func] # l2 or dot product function
75
76 def setup_faiss(self):
77 if not self.val_file:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected