Views
No views yet
1import torch
2import os
3import random
4import numpy as np
5import json
6import re
7
8from torch import Tensor
9from transformers import AutoTokenizer, AutoModelForCausalLM
10
11from prompts_synthesis import get_create_classify_data_prompt
12from utils import fix_common_json_errors_and_loads
13
14
15LLAMA3_PROMPT = """
16{prompt} [/INST]
17""".strip("\n")
18
19# Each query must come with a one-sentence instruction that describes the task
20tasks = [
21 'Identify the intended age group for educational technology products.',
22 'Classify businesses based on their operational hours.'
23]
24language = 'English'
25
26prompts = [LLAMA3_PROMPT.format(prompt=get_create_classify_data_prompt(task=task, language=language)[1]['content']) for task in tasks]
27
28tokenizer = AutoTokenizer.from_pretrained('Haon-Chen/speed-synthesis-7b-senior')
29model = AutoModelForCausalLM.from_pretrained('Haon-Chen/speed-synthesis-7b-senior')
30model.to("cuda:0")
31model.eval()
32tokenizer.pad_token = tokenizer.pad_token or tokenizer.eos_token
33tokenizer.padding_side = "left"
34tokenizer.truncation_side = "left"
35
36with torch.inference_mode():
37 # Tokenize the input texts
38 encodes = tokenizer(prompts, padding="longest", add_special_tokens=True, return_tensors="pt")
39 input_ids = encodes.input_ids.to(model.device)
40 attention_mask = encodes.attention_mask.to(model.device)
41
42 # Set the generation parameters
43 GEN_CONFIG = {"do_sample":True, "temperature": 1.0, "top_p": 1.0, "max_new_tokens": 800}
44 output = model.generate(
45 input_ids=input_ids,
46 attention_mask=attention_mask,
47 pad_token_id = tokenizer.eos_token_id,
48 **GEN_CONFIG
49 )
50output_texts = tokenizer.batch_decode(output, skip_special_tokens=True, clean_up_tokenization_spaces=False)
51batch_results = []
52for i in range(len(output_texts)):
53 batch_results.append(output_texts[i][len(prompts[i]):].strip(' '))
54
55# Format outputs
56bad_cnt=0
57outputs = []
58for i, result in enumerate(batch_results):
59 try:
60 output = fix_common_json_errors_and_loads(result)
61 user_query = output.get("input_text", "")
62 positive_document = output.get("label", "")
63 hard_negative_document = output.get("misleading_label", "")
64 except:
65 bad_cnt+=1
66 continue
67 out_data = {
68 "query": user_query,
69 "positives": [positive_document],
70 "negatives": [hard_negative_document],
71 "language": "English",
72 "task_definition": tasks[i],
73 }
74 outputs.append(out_data)
75print(bad_cnt)
76print(outputs)1@article{chen2024little,
2 title={Little Giants: Synthesizing High-Quality Embedding Data at Scale},
3 author={Chen, Haonan and Wang, Liang and Yang, Nan and Zhu, Yutao and Zhao, Ziliang and Wei, Furu and Dou, Zhicheng},
4 journal={arXiv preprint arXiv:2410.18634},
5 year={2024}
6}