Views
No views yet
1# %% ByteETM Inference (바이트 기반 추론)
2import torch
3from transformers import AutoModelForCausalLM
4
5# 1️⃣ 모델 로드
6repo_id = "idah4/byteetm-korean-tiny"
7device = "cuda" if torch.cuda.is_available() else "cpu"
8
9model = AutoModelForCausalLM.from_pretrained(
10 repo_id,
11 trust_remote_code=True
12).to(device).eval()
13
14# 2️⃣ 바이트 기반 인코더 / 디코더
15def encode_bytes(text: str):
16 return torch.tensor([[b for b in text.encode("utf-8")]], dtype=torch.long, device=device)
17
18def decode_bytes(ids: torch.Tensor):
19 seq = [i for i in ids.tolist() if 0 <= i < 256]
20 return bytes(seq).decode("utf-8", errors="ignore")
21
22# 3️⃣ 텍스트 생성 함수
23@torch.no_grad()
24def generate_text(prompt: str, max_new_tokens=200, temperature=0.8, top_k=200):
25 input_ids = encode_bytes(prompt)
26 out = model.generate(
27 input_ids,
28 max_new_tokens=max_new_tokens,
29 temperature=temperature,
30 top_k=top_k
31 )
32 return decode_bytes(out[0])
33
34# 4️⃣ 시연
35prompt = "오늘은 날씨가 좋아서"
36print(generate_text(prompt, max_new_tokens=150, temperature=0.9, top_k=150))
37