Views
No views yet
1import os
2from huggingface_hub import snapshot_download
3
4# 경로 설정
5cache_dir = "--Enter Your Desired File Direction--"
6tmp_dir = "--Enter Your Desired File Direction--"
7local_dir = "--Enter Your Desired File Direction--"
8repo_id = "ChaeSJ/llama-3.1-8b-finetuned"
9
10os.environ["TRANSFORMERS_CACHE"] = cache_dir
11os.environ["HF_HOME"] = cache_dir
12os.makedirs(cache_dir, exist_ok=True)
13os.makedirs(tmp_dir, exist_ok=True)
14os.environ["TMPDIR"] = tmp_dir
15
16print(f"Downloading snapshot of {repo_id} ...")
17snapshot_download(
18 repo_id=repo_id,
19 local_dir=local_dir,
20 local_dir_use_symlinks=False,
21 cache_dir=cache_dir,
22 resume_download=True
23)
24print(f" Downloaded all files to: {local_dir}")
251
2import os
3import json
4import re
5import torch
6from tqdm import tqdm
7from transformers import AutoTokenizer, AutoModelForCausalLM
8
9REPO_ID = "ChaeSJ/llama-3.1-8b-finetuned"
10SUBFOLDER = "llama_3_1_8b_finetuned"
11CACHE_DIR = "--Enter Your Desired File Direction--"
12TMP_DIR = "--Enter Your Desired File Direction--"
13
14os.environ["TRANSFORMERS_VERBOSITY"] = "error"
15os.environ["TRANSFORMERS_CACHE"] = CACHE_DIR
16os.environ["HF_HOME"] = CACHE_DIR
17os.makedirs(CACHE_DIR, exist_ok=True)
18os.makedirs(TMP_DIR, exist_ok=True)
19os.environ["TMPDIR"] = TMP_DIR
20
21## 로컬 입출력 경로(프롬프트/결과물)
22input_path = "--Enter Your Desired File Direction--/selected_prompts_c++.json"
23json_output_path = "--Enter Your Desired File Direction--/llama_3_1_8b_demo_generated.json"
24cpp_output_dir = "--Enter Your Desired File Direction--/extracted_cpp"
25
26## 코드 블록에서 C/C++ 추출
27def extract_cpp_code(text: str) -> str:
28 m = re.findall(r"```(?:cpp|c\+\+)\n(.*?)```", text, re.DOTALL | re.IGNORECASE)
29 if m:
30 return m[0].strip()
31 m = re.findall(r"```c\n(.*?)```", text, re.DOTALL | re.IGNORECASE)
32 if m:
33 return m[0].strip()
34 return text.strip()
35
36def apply_template(tokenizer, prompt_text: str) -> str:
37 return tokenizer.apply_chat_template(
38 [{"role": "user", "content": prompt_text}],
39 tokenize=False,
40 add_generation_prompt=True,
41 )
42
43## 토크나이저 & 모델: 허깅페이스에서 직접 로드
44print(">> Loading tokenizer from Hugging Face...")
45tokenizer = AutoTokenizer.from_pretrained(
46 REPO_ID,
47 subfolder=SUBFOLDER,
48 cache_dir=CACHE_DIR,
49 use_fast=True,
50 padding_side="right",
51)
52
53if tokenizer.pad_token is None:
54 tokenizer.add_special_tokens({"pad_token": tokenizer.eos_token})
55 print(">> Added pad_token as eos_token")
56
57print(">> Loading model from Hugging Face (full finetuned weights)...")
58model = AutoModelForCausalLM.from_pretrained(
59 REPO_ID,
60 subfolder=SUBFOLDER,
61 cache_dir=CACHE_DIR,
62 torch_dtype=torch.bfloat16,
63 device_map="auto",
64 low_cpu_mem_usage=True,
65)
66
67if model.config.pad_token_id is None:
68 model.config.pad_token_id = tokenizer.pad_token_id
69
70model.eval()
71
72## 프롬프트 로드
73print(">> Loading prompts...")
74with open(input_path, "r", encoding="utf-8") as f:
75 prompts_data = json.load(f)
76
77## 생성 & 저장
78print(">> Generating code...")
79os.makedirs(os.path.dirname(json_output_path), exist_ok=True)
80os.makedirs(cpp_output_dir, exist_ok=True)
81
82results = []
83
84@torch.inference_mode()
85def generate_one(prompt_text: str, max_input_len=1536, max_new_tokens=512) -> str:
86 templated = apply_template(tokenizer, prompt_text)
87
88 inputs = tokenizer(
89 templated,
90 return_tensors="pt",
91 padding=False,
92 truncation=True,
93 max_length=max_input_len,
94 )
95
96 for k in inputs:
97 inputs[k] = inputs[k].to(model.device)
98
99 out = model.generate(
100 **inputs,
101 max_new_tokens=max_new_tokens,
102 pad_token_id=tokenizer.pad_token_id,
103 eos_token_id=tokenizer.eos_token_id,
104
105 do_sample=False,
106 num_beams=1,
107 temperature=0.0,
108 top_p=1.0,
109 top_k=0,
110 repetition_penalty=1.0,
111 no_repeat_ngram_size=0,
112
113 use_cache=True,
114 return_dict_in_generate=True,
115 )
116
117 input_len = inputs["input_ids"].shape[1]
118 gen_ids = out.sequences[0][input_len:]
119 return tokenizer.decode(gen_ids, skip_special_tokens=True)
120
121for idx, item in enumerate(tqdm(prompts_data, desc="Generating")):
122 prompt_text = item.get("nl_prompt", "")
123 prompt_id = str(item.get("Prompt ID", "")).strip() or None
124
125 generated_text = generate_one(prompt_text, max_input_len=1536, max_new_tokens=512)
126
127 results.append({
128 "Prompt ID": prompt_id if prompt_id is not None else f"{idx+1}",
129 "nl_prompt": prompt_text,
130 "generated_code": generated_text
131 })
132
133## JSON 저장
134print(f">> Saving generated results to {json_output_path}")
135with open(json_output_path, "w", encoding="utf-8") as f:
136 json.dump(results, f, indent=2, ensure_ascii=False)
137
138## C/C++ 코드 추출 및 파일 저장
139print(f">> Extracting cpp code to directory: {cpp_output_dir}")
140saved = 0
141for i, item in enumerate(results):
142 raw_code = item.get("generated_code", "")
143 code = extract_cpp_code(raw_code)
144
145 pid = item.get("Prompt ID")
146 if pid and pid != "unknown":
147 safe_pid = re.sub(r"[^a-zA-Z0-9_\-\.]+", "_", str(pid))[:64]
148 filename = f"Llama_3.1_8b_demo_{safe_pid}.cpp"
149 else:
150 filename = f"Llama_3.1_8b_demo_{i+1:04d}.cpp"
151
152 filepath = os.path.join(cpp_output_dir, filename)
153 with open(filepath, "w", encoding="utf-8") as w:
154 w.write(code)
155 saved += 1
156
157print(f">> Extraction complete. {saved} files saved to {cpp_output_dir}")
158