MCPcopy Create free account
hub / github.com/FlyingFeather/DEA-SQL / gen_embed_and_save

Function gen_embed_and_save

fewshot/processor.py:49–110  ·  view source on GitHub ↗
(input_file, version, key, model='transformer')

Source from the content-addressed store, hash-verified

47
48
49def gen_embed_and_save(input_file, version, key, model='transformer'):
50
51 # validate keys
52 assert key in ['query', 'pred', 'masked_query']
53 assert model in ['openai', 'transformer']
54
55 print(f'Start embedding...')
56 print(f'data: {input_file}')
57 print(f'key: {key}')
58 print(f'model: {model}')
59
60 # Read and process the input data
61 input_file = input_file
62 data = {}
63 queries = []
64 samples = []
65 with open(input_file, 'r') as f:
66 for line in f:
67 sample = json.loads(line) # Parse JSON text
68 query = sample[key] # Extract key to embed
69 queries.append(query)
70 samples.append(sample)
71
72 # Convert queries to vectors in a batch
73 query_vectors = []
74 if model == 'openai':
75 query_vectors = []
76 max_batch_size = 16
77 for i in range(0, len(queries), max_batch_size):
78 sub_queries = queries[i:i + max_batch_size]
79 sub_query_vectors = embed_by_openai(sub_queries)
80 query_vectors += sub_query_vectors
81 elif model == 'transformer':
82 query_vectors = None
83 max_batch_size = 200
84 print(f'batch size: {max_batch_size}')
85 for i in range(0, len(queries), max_batch_size):
86 print(f'current batch {i}:{i + max_batch_size}')
87 sub_queries = queries[i:i + max_batch_size]
88 sub_query_vectors = embed_by_transformers(sub_queries)
89 if query_vectors is None:
90 query_vectors = sub_query_vectors
91 else:
92 query_vectors = np.vstack((query_vectors, sub_query_vectors))
93
94 query_vectors = np.array(query_vectors)
95 print(query_vectors.shape)
96
97 # generate key-value pair {embedding: sample}
98 for query_vector, sample in zip(query_vectors, samples):
99 data[tuple(query_vector)] = sample
100
101 # Save the results to file
102 output_file = f'{key}_embed_{model}.pkl'
103 output_dir = f'fewshot/index_{version}/'
104
105 if not os.path.exists(output_dir):
106 os.makedirs(output_dir)

Callers 1

processor.pyFile · 0.85

Calls 3

embed_by_openaiFunction · 0.90
embed_by_transformersFunction · 0.90
dumpMethod · 0.80

Tested by

no test coverage detected