Views
No views yet
1pip install transformers
2pip install peft
3pip install bitsandbytes
4pip install acceleratebash merge.sh1from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
2from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
3import torch
4
5
6# load model and tokenizer
7
8tokenizer = AutoTokenizer.from_pretrained("CodeLlama-70B_for_NTR/Epoch_1/-merged", use_auth_token=True)
9
10nf4_config = BitsAndBytesConfig(
11load_in_4bit=True,
12bnb_4bit_quant_type="nf4",
13bnb_4bit_use_double_quant=True,
14bnb_4bit_compute_dtype=torch.bfloat16
15)
16
17model = AutoModelForCausalLM.from_pretrained(
18 "CodeLlama-70B_for_NTR/Epoch_1/-merged",
19 quantization_config=nf4_config,
20 device_map='auto'
21)
22
23model = prepare_model_for_kbit_training(model)
24
25lora_config = LoraConfig(
26 r=16,
27 lora_alpha=32,
28 lora_dropout=0.05,
29 bias="none",
30 task_type="CAUSAL_LM",
31 target_modules = ["q_proj", "k_proj", "v_proj", "o_proj"]
32)
33
34model = get_peft_model(model, lora_config)
35
36
37# a bug-fix pairs
38
39buggy_code = "
40 public MultiplePiePlot(CategoryDataset dataset){
41 super();
42// bug_start
43 this.dataset=dataset;
44// bug_end
45 PiePlot piePlot=new PiePlot(null);
46 this.pieChart=new JFreeChart(piePlot);
47 this.pieChart.removeLegend();
48 this.dataExtractOrder=TableOrder.BY_COLUMN;
49 this.pieChart.setBackgroundPaint(null);
50 TextTitle seriesTitle=new TextTitle("Series Title",new Font("SansSerif",Font.BOLD,12));
51 seriesTitle.setPosition(RectangleEdge.BOTTOM);
52 this.pieChart.setTitle(seriesTitle);
53 this.aggregatedItemsKey="Other";
54 this.aggregatedItemsPaint=Color.lightGray;
55 this.sectionPaints=new HashMap();
56 }
57"
58
59repair_template = "OtherTemplate"
60
61fixed_code = "
62// fix_start
63 setDataset(dataset);
64// fix_end
65"
66
67# model inference
68
69B_INST, E_INST = "[INST]", "[/INST]"
70input_text = tokenizer.bos_token + B_INST +'\n[bug_function]\n' + buggy_code + '\n[fix_template]\n' + repair_template + '\n[fix_code]\n' + E_INST
71input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(0)
72
73eos_id = tokenizer.convert_tokens_to_ids(tokenizer.eos_token)
74generated_ids = model.generate(
75 input_ids=input_ids,
76 max_new_tokens=256,
77 num_beams=10,
78 num_return_sequences=10,
79 early_stopping=True,
80 pad_token_id=eos_id,
81 eos_token_id=eos_id
82)
83
84for generated_id in generated_ids:
85 generated_text = tokenizer.decode(generated_id, skip_special_tokens=False)
86 patch = generated_text.split(E_INST)[1]
87 patch = patch.replace(tokenizer.eos_token,'')
88 print(patch)
89
90