Skip to content

Commit 58b23c3

Browse files
committed
refactor: KULLM3 4bit 양자화 및 배치 추론 기반 LLM 모듈로 summarizer.py 리팩토링
1 parent dfff4ad commit 58b23c3

1 file changed

Lines changed: 79 additions & 11 deletions

File tree

app/services/summarizer.py

Lines changed: 79 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,20 +1,48 @@
11
# 비즈니스 로직 / AI 추론 모듈
22

3-
from transformers import AutoModelForCausalLM, AutoTokenizer
43
import torch
4+
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
5+
import time
6+
from typing import List
57

6-
# 모델과 토크나이저를 전역으로 로드 (최초 1회만)
7-
tokenizer = AutoTokenizer.from_pretrained("nlpai-lab/KULLM3")
8-
model = AutoModelForCausalLM.from_pretrained("nlpai-lab/KULLM3")
8+
# 1. 모델 불러오기 및 4bit 양자화 설정
9+
model_id = "nlpai-lab/KULLM3"
10+
bnb_config = BitsAndBytesConfig(
11+
load_in_4bit=True,
12+
bnb_4bit_use_double_quant=True,
13+
bnb_4bit_quant_type="nf4",
14+
bnb_4bit_compute_dtype=torch.bfloat16
15+
)
916

10-
def generate_content(prompt: str) -> str:
11-
inputs = tokenizer(prompt, return_tensors="pt")
12-
with torch.no_grad():
13-
outputs = model.generate(**inputs, max_new_tokens=256)
14-
result = tokenizer.decode(outputs[0], skip_special_tokens=True)
15-
# 프롬프트 부분 제거 (모델에 따라 필요)
16-
return result[len(prompt):].strip() if result.startswith(prompt) else result
17+
tokenizer = AutoTokenizer.from_pretrained(model_id)
1718

19+
if torch.cuda.is_available():
20+
torch.cuda.empty_cache()
21+
torch.cuda.reset_peak_memory_stats()
22+
print(f"Initial VRAM usage: {torch.cuda.memory_allocated() / (1024**3):.2f} GB")
23+
24+
start_load_time = time.time()
25+
model = AutoModelForCausalLM.from_pretrained(
26+
model_id,
27+
quantization_config=bnb_config,
28+
device_map="auto"
29+
)
30+
end_load_time = time.time()
31+
32+
print(f"\nModel loaded in {end_load_time - start_load_time:.2f} seconds")
33+
34+
if torch.cuda.is_available():
35+
initial_vram_after_load = torch.cuda.memory_allocated()
36+
peak_vram_after_load = torch.cuda.max_memory_allocated()
37+
print(f"VRAM allocated after model load: {initial_vram_after_load / (1024**3):.2f} GB")
38+
print(f"Peak VRAM used during model load: {peak_vram_after_load / (1024**3):.2f} GB")
39+
40+
# 2. LLM 채팅 프롬프트 포맷
41+
42+
def build_chat_prompt(prompt: str):
43+
return f"<s>[INST] {prompt.strip()} [/INST]"
44+
45+
# 3. 프롬프트 생성 함수 (기존 유지)
1846
def build_transform_prompt(title: str, content: str, level: str) -> str:
1947
base = f"다음 뉴스 제목과 본문을 사용자의 이해 수준에 맞게 다시 써줘.\n\n뉴스 제목: {title}\n뉴스 본문: {content}\n"
2048
if level == "상":
@@ -27,3 +55,43 @@ def build_transform_prompt(title: str, content: str, level: str) -> str:
2755

2856
def build_summary_prompt(title: str, content: str) -> str:
2957
return f"다음 뉴스 제목과 본문을 한문장으로 간단히 요약해줘.\n\n뉴스 제목: {title}\n뉴스 본문: {content}"
58+
59+
# 4. 배치 추론 함수
60+
def kullm_batch_generate(prompts: List[str], max_new_tokens=512):
61+
chat_prompts = [build_chat_prompt(p) for p in prompts]
62+
if torch.cuda.is_available():
63+
torch.cuda.reset_peak_memory_stats()
64+
inputs = tokenizer(chat_prompts, return_tensors="pt", padding=True).to(model.device)
65+
input_ids = inputs.input_ids
66+
attention_mask = inputs.attention_mask
67+
start_infer_time = time.time()
68+
output = model.generate(
69+
input_ids=input_ids,
70+
attention_mask=attention_mask,
71+
max_new_tokens=max_new_tokens,
72+
do_sample=True,
73+
temperature=0.2,
74+
top_p=0.2,
75+
pad_token_id=tokenizer.eos_token_id
76+
)
77+
end_infer_time = time.time()
78+
generation_time = end_infer_time - start_infer_time
79+
decoded_results = []
80+
generated_tokens_list = []
81+
for i in range(len(prompts)):
82+
original_input_len = (input_ids[i] != tokenizer.pad_token_id).sum().item()
83+
generated_tokens = output[i].shape[0] - original_input_len
84+
generated_tokens_list.append(generated_tokens)
85+
result_text = tokenizer.decode(output[i], skip_special_tokens=True)
86+
decoded_results.append(result_text.split('[/INST]')[-1].strip())
87+
current_vram = 0
88+
peak_vram = 0
89+
if torch.cuda.is_available():
90+
current_vram = torch.cuda.memory_allocated()
91+
peak_vram = torch.cuda.max_memory_allocated()
92+
return decoded_results, generation_time, generated_tokens_list, current_vram, peak_vram
93+
94+
# 5. 단일 프롬프트용 generate_content 함수
95+
def generate_content(prompt: str, max_new_tokens=512) -> str:
96+
results, _, _, _, _ = kullm_batch_generate([prompt], max_new_tokens=max_new_tokens)
97+
return results[0]

0 commit comments

Comments
 (0)