(input_file, version, key, model='transformer')
| 47 | |
| 48 | |
| 49 | def 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) |
no test coverage detected