11# 비즈니스 로직 / AI 추론 모듈
22
3- from transformers import AutoModelForCausalLM , AutoTokenizer
43import 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"\n Model 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. 프롬프트 생성 함수 (기존 유지)
1846def 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
2856def 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