13 KiB
๋๊ท๋ชจ ์ธ์ด ๋ชจ๋ธ๋ก ์์ฑํ๊ธฐ generation-with-llms
LLM ๋๋ ๋๊ท๋ชจ ์ธ์ด ๋ชจ๋ธ์ ํ ์คํธ ์์ฑ์ ํต์ฌ ๊ตฌ์ฑ ์์์ ๋๋ค. ๊ฐ๋จํ ๋งํ๋ฉด, ์ฃผ์ด์ง ์ ๋ ฅ ํ ์คํธ์ ๋ํ ๋ค์ ๋จ์ด(์ ํํ๊ฒ๋ ํ ํฐ)๋ฅผ ์์ธกํ๊ธฐ ์ํด ํ๋ จ๋ ๋๊ท๋ชจ ์ฌ์ ํ๋ จ ๋ณํ๊ธฐ ๋ชจ๋ธ๋ก ๊ตฌ์ฑ๋ฉ๋๋ค. ํ ํฐ์ ํ ๋ฒ์ ํ๋์ฉ ์์ธกํ๊ธฐ ๋๋ฌธ์ ์๋ก์ด ๋ฌธ์ฅ์ ์์ฑํ๋ ค๋ฉด ๋ชจ๋ธ์ ํธ์ถํ๋ ๊ฒ ์ธ์ ๋ ๋ณต์กํ ์์ ์ ์ํํด์ผ ํฉ๋๋ค. ์ฆ, ์๊ธฐํ๊ท ์์ฑ์ ์ํํด์ผ ํฉ๋๋ค.
์๊ธฐํ๊ท ์์ฑ์ ๋ช ๊ฐ์ ์ด๊ธฐ ์
๋ ฅ๊ฐ์ ์ ๊ณตํ ํ, ๊ทธ ์ถ๋ ฅ์ ๋ค์ ๋ชจ๋ธ์ ์
๋ ฅ์ผ๋ก ์ฌ์ฉํ์ฌ ๋ฐ๋ณต์ ์ผ๋ก ํธ์ถํ๋ ์ถ๋ก ๊ณผ์ ์
๋๋ค. ๐ค Transformers์์๋ [~generation.GenerationMixin.generate] ๋ฉ์๋๊ฐ ์ด ์ญํ ์ ํ๋ฉฐ, ์ด๋ ์์ฑ ๊ธฐ๋ฅ์ ๊ฐ์ง ๋ชจ๋ ๋ชจ๋ธ์์ ์ฌ์ฉ ๊ฐ๋ฅํฉ๋๋ค.
์ด ํํ ๋ฆฌ์ผ์์๋ ๋ค์ ๋ด์ฉ์ ๋ค๋ฃจ๊ฒ ๋ฉ๋๋ค:
- LLM์ผ๋ก ํ ์คํธ ์์ฑ
- ์ผ๋ฐ์ ์ผ๋ก ๋ฐ์ํ๋ ๋ฌธ์ ํด๊ฒฐ
- LLM์ ์ต๋ํ ํ์ฉํ๊ธฐ ์ํ ๋ค์ ๋จ๊ณ
์์ํ๊ธฐ ์ ์ ํ์ํ ๋ชจ๋ ๋ผ์ด๋ธ๋ฌ๋ฆฌ๊ฐ ์ค์น๋์ด ์๋์ง ํ์ธํ์ธ์:
pip install transformers bitsandbytes>=0.39.0 -q
ํ ์คํธ ์์ฑ generate-text
์ธ๊ณผ์ ์ธ์ด ๋ชจ๋ธ๋ง(causal language modeling)์ ๋ชฉ์ ์ผ๋ก ํ์ต๋ ์ธ์ด ๋ชจ๋ธ์ ์ผ๋ จ์ ํ ์คํธ ํ ํฐ์ ์ ๋ ฅ์ผ๋ก ์ฌ์ฉํ๊ณ , ๊ทธ ๊ฒฐ๊ณผ๋ก ๋ค์ ํ ํฐ์ด ๋์ฌ ํ๋ฅ ๋ถํฌ๋ฅผ ์ ๊ณตํฉ๋๋ค.
LLM๊ณผ ์๊ธฐํ๊ท ์์ฑ์ ํจ๊ป ์ฌ์ฉํ ๋ ํต์ฌ์ ์ธ ๋ถ๋ถ์ ์ด ํ๋ฅ ๋ถํฌ๋ก๋ถํฐ ๋ค์ ํ ํฐ์ ์ด๋ป๊ฒ ๊ณ ๋ฅผ ๊ฒ์ธ์ง์ ๋๋ค. ๋ค์ ๋ฐ๋ณต ๊ณผ์ ์ ์ฌ์ฉ๋ ํ ํฐ์ ๊ฒฐ์ ํ๋ ํ, ์ด๋ ํ ๋ฐฉ๋ฒ๋ ๊ฐ๋ฅํฉ๋๋ค. ํ๋ฅ ๋ถํฌ์์ ๊ฐ์ฅ ๊ฐ๋ฅ์ฑ์ด ๋์ ํ ํฐ์ ์ ํํ๋ ๊ฒ์ฒ๋ผ ๊ฐ๋จํ ์๋ ์๊ณ , ๊ฒฐ๊ณผ ๋ถํฌ์์ ์ํ๋งํ๊ธฐ ์ ์ ์์ญ ๊ฐ์ง ๋ณํ์ ์ ์ฉํ๋ ๊ฒ์ฒ๋ผ ๋ณต์กํ ์๋ ์์ต๋๋ค.
์์์ ์ค๋ช ํ ๊ณผ์ ์ ์ด๋ค ์ข ๋ฃ ์กฐ๊ฑด์ด ์ถฉ์กฑ๋ ๋๊น์ง ๋ฐ๋ณต์ ์ผ๋ก ์ํ๋ฉ๋๋ค. ๋ชจ๋ธ์ด ์ํ์ค์ ๋(EOS ํ ํฐ)์ ์ถ๋ ฅํ ๋๊น์ง๋ฅผ ์ข ๋ฃ ์กฐ๊ฑด์ผ๋ก ํ๋ ๊ฒ์ด ์ด์์ ์ ๋๋ค. ๊ทธ๋ ์ง ์์ ๊ฒฝ์ฐ์๋ ๋ฏธ๋ฆฌ ์ ์๋ ์ต๋ ๊ธธ์ด์ ๋๋ฌํ์ ๋ ์์ฑ์ด ์ค๋จ๋ฉ๋๋ค.
๋ชจ๋ธ์ด ์์๋๋ก ๋์ํ๊ธฐ ์ํด์ ํ ํฐ ์ ํ ๋จ๊ณ์ ์ ์ง ์กฐ๊ฑด์ ์ฌ๋ฐ๋ฅด๊ฒ ์ค์ ํ๋ ๊ฒ์ด ์ค์ํฉ๋๋ค. ์ด๋ฌํ ์ด์ ๋ก, ๊ฐ ๋ชจ๋ธ์๋ ๊ธฐ๋ณธ ์์ฑ ์ค์ ์ด ์ ์ ์๋ [~generation.GenerationConfig] ํ์ผ์ด ํจ๊ป ์ ๊ณต๋ฉ๋๋ค.
์ฝ๋๋ฅผ ํ์ธํด๋ด ์๋ค!
๊ธฐ๋ณธ LLM ์ฌ์ฉ์ ๊ด์ฌ์ด ์๋ค๋ฉด, ์ฐ๋ฆฌ์ Pipeline ์ธํฐํ์ด์ค๋ก ์์ํ๋ ๊ฒ์ ์ถ์ฒํฉ๋๋ค. ๊ทธ๋ฌ๋ LLM์ ์์ํ๋ ํ ํฐ ์ ํ ๋จ๊ณ์์์ ๋ฏธ์ธํ ์ ์ด์ ๊ฐ์ ๊ณ ๊ธ ๊ธฐ๋ฅ๋ค์ ์ข
์ข
ํ์๋ก ํฉ๋๋ค. ์ด๋ฌํ ์์
์ [~generation.GenerationMixin.generate]๋ฅผ ํตํด ๊ฐ์ฅ ์ ์ํ๋ ์ ์์ต๋๋ค. LLM์ ์ด์ฉํ ์๊ธฐํ๊ท ์์ฑ์ ์์์ ๋ง์ด ์๋ชจํ๋ฏ๋ก, ์ ์ ํ ์ฒ๋ฆฌ๋์ ์ํด GPU์์ ์คํ๋์ด์ผ ํฉ๋๋ค.
๋จผ์ , ๋ชจ๋ธ์ ๋ถ๋ฌ์ค์ธ์.
>>> from transformers import AutoModelForCausalLM
>>> model = AutoModelForCausalLM.from_pretrained(
... "mistralai/Mistral-7B-v0.1", device_map="auto", load_in_4bit=True
... )
from_pretrained ํจ์๋ฅผ ํธ์ถํ ๋ 2๊ฐ์ ํ๋๊ทธ๋ฅผ ์ฃผ๋ชฉํ์ธ์:
device_map์ ๋ชจ๋ธ์ด GPU๋ก ์ด๋๋๋๋ก ํฉ๋๋ค.load_in_4bit๋ ๋ฆฌ์์ค ์๊ตฌ ์ฌํญ์ ํฌ๊ฒ ์ค์ด๊ธฐ ์ํด 4๋นํธ ๋์ ์์ํ๋ฅผ ์ ์ฉํฉ๋๋ค.
์ด ์ธ์๋ ๋ชจ๋ธ์ ์ด๊ธฐํํ๋ ๋ค์ํ ๋ฐฉ๋ฒ์ด ์์ง๋ง, LLM์ ์ฒ์ ์์ํ ๋ ์ด ์ค์ ์ ์ถ์ฒํฉ๋๋ค.
์ด์ด์ ํ ์คํธ ์ ๋ ฅ์ ํ ํฌ๋์ด์ ์ผ๋ก ์ ์ฒ๋ฆฌํ์ธ์.
>>> from transformers import AutoTokenizer
>>> import torch
>>> tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-v0.1")
>>> device = "cuda" if torch.cuda.is_available() else "cpu"
>>> model_inputs = tokenizer(["A list of colors: red, blue"], return_tensors="pt").to(device)
model_inputs ๋ณ์์๋ ํ ํฐํ๋ ํ
์คํธ ์
๋ ฅ๊ณผ ํจ๊ป ์ดํ
์
๋ง์คํฌ๊ฐ ๋ค์ด ์์ต๋๋ค. [~generation.GenerationMixin.generate]๋ ์ดํ
์
๋ง์คํฌ๊ฐ ์ ๊ณต๋์ง ์์์ ๊ฒฝ์ฐ์๋ ์ด๋ฅผ ์ถ๋ก ํ๋ ค๊ณ ๋
ธ๋ ฅํ์ง๋ง, ์ต์์ ์ฑ๋ฅ์ ์ํด์๋ ๊ฐ๋ฅํ๋ฉด ์ดํ
์
๋ง์คํฌ๋ฅผ ์ ๋ฌํ๋ ๊ฒ์ ๊ถ์ฅํฉ๋๋ค.
๋ง์ง๋ง์ผ๋ก [~generation.GenerationMixin.generate] ๋ฉ์๋๋ฅผ ํธ์ถํด ์์ฑ๋ ํ ํฐ์ ์ป์ ํ, ์ด๋ฅผ ์ถ๋ ฅํ๊ธฐ ์ ์ ํ
์คํธ ํํ๋ก ๋ณํํ์ธ์.
>>> generated_ids = model.generate(**model_inputs)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'A list of colors: red, blue, green, yellow, black, white, and brown'
์ด๊ฒ ์ ๋ถ์ ๋๋ค! ๋ช ์ค์ ์ฝ๋๋ง์ผ๋ก LLM์ ๋ฅ๋ ฅ์ ํ์ฉํ ์ ์๊ฒ ๋์์ต๋๋ค.
์ผ๋ฐ์ ์ผ๋ก ๋ฐ์ํ๋ ๋ฌธ์ common-pitfalls
์์ฑ ์ ๋ต์ด ๋ง๊ณ , ๊ธฐ๋ณธ๊ฐ์ด ํญ์ ์ฌ์ฉ ์ฌ๋ก์ ์ ํฉํ์ง ์์ ์ ์์ต๋๋ค. ์ถ๋ ฅ์ด ์์๊ณผ ๋ค๋ฅผ ๋ ํํ ๋ฐ์ํ๋ ๋ฌธ์ ์ ์ด๋ฅผ ํด๊ฒฐํ๋ ๋ฐฉ๋ฒ์ ๋ํ ๋ชฉ๋ก์ ๋ง๋ค์์ต๋๋ค.
>>> from transformers import AutoModelForCausalLM, AutoTokenizer
>>> tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-v0.1")
>>> tokenizer.pad_token = tokenizer.eos_token # Mistral has no pad token by default
>>> model = AutoModelForCausalLM.from_pretrained(
... "mistralai/Mistral-7B-v0.1", device_map="auto", load_in_4bit=True
... )
์์ฑ๋ ์ถ๋ ฅ์ด ๋๋ฌด ์งง๊ฑฐ๋ ๊ธธ๋ค generated-output-is-too-shortlong
[~generation.GenerationConfig] ํ์ผ์์ ๋ณ๋๋ก ์ง์ ํ์ง ์์ผ๋ฉด, generate๋ ๊ธฐ๋ณธ์ ์ผ๋ก ์ต๋ 20๊ฐ์ ํ ํฐ์ ๋ฐํํฉ๋๋ค. generate ํธ์ถ์์ max_new_tokens์ ์๋์ผ๋ก ์ค์ ํ์ฌ ๋ฐํํ ์ ์๋ ์ ํ ํฐ์ ์ต๋ ์๋ฅผ ์ค์ ํ๋ ๊ฒ์ด ์ข์ต๋๋ค. LLM(์ ํํ๊ฒ๋ ๋์ฝ๋ ์ ์ฉ ๋ชจ๋ธ)์ ์
๋ ฅ ํ๋กฌํํธ๋ ์ถ๋ ฅ์ ์ผ๋ถ๋ก ๋ฐํํฉ๋๋ค.
>>> model_inputs = tokenizer(["A sequence of numbers: 1, 2"], return_tensors="pt").to("cuda")
>>> # By default, the output will contain up to 20 tokens
>>> generated_ids = model.generate(**model_inputs, pad_token_id=tokenizer.eos_token_id)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'A sequence of numbers: 1, 2, 3, 4, 5'
>>> # Setting `max_new_tokens` allows you to control the maximum length
>>> generated_ids = model.generate(**model_inputs, pad_token_id=tokenizer.eos_token_id, max_new_tokens=50)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'A sequence of numbers: 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16,'
์๋ชป๋ ์์ฑ ๋ชจ๋ incorrect-generation-mode
๊ธฐ๋ณธ์ ์ผ๋ก [~generation.GenerationConfig] ํ์ผ์์ ๋ณ๋๋ก ์ง์ ํ์ง ์์ผ๋ฉด, generate๋ ๊ฐ ๋ฐ๋ณต์์ ๊ฐ์ฅ ํ๋ฅ ์ด ๋์ ํ ํฐ์ ์ ํํฉ๋๋ค(๊ทธ๋ฆฌ๋ ๋์ฝ๋ฉ). ํ๋ ค๋ ์์
์ ๋ฐ๋ผ ์ด ๋ฐฉ๋ฒ์ ๋ฐ๋์งํ์ง ์์ ์ ์์ต๋๋ค. ์๋ฅผ ๋ค์ด, ์ฑ๋ด์ด๋ ์์ธ์ด ์์ฑ๊ณผ ๊ฐ์ ์ฐฝ์์ ์ธ ์์
์ ์ํ๋ง์ด ์ ํฉํ ์ ์์ต๋๋ค. ๋ฐ๋ฉด, ์ค๋์ค๋ฅผ ํ
์คํธ๋ก ๋ณํํ๊ฑฐ๋ ๋ฒ์ญ๊ณผ ๊ฐ์ ์
๋ ฅ ๊ธฐ๋ฐ ์์
์ ๊ทธ๋ฆฌ๋ ๋์ฝ๋ฉ์ด ๋ ์ ํฉํ ์ ์์ต๋๋ค. do_sample=True๋ก ์ํ๋ง์ ํ์ฑํํ ์ ์์ผ๋ฉฐ, ์ด ์ฃผ์ ์ ๋ํ ์์ธํ ๋ด์ฉ์ ์ด ๋ธ๋ก๊ทธ ํฌ์คํธ์์ ๋ณผ ์ ์์ต๋๋ค.
>>> # Set seed or reproducibility -- you don't need this unless you want full reproducibility
>>> from transformers import set_seed
>>> set_seed(0)
>>> model_inputs = tokenizer(["I am a cat."], return_tensors="pt").to("cuda")
>>> # LLM + greedy decoding = repetitive, boring output
>>> generated_ids = model.generate(**model_inputs)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'I am a cat. I am a cat. I am a cat. I am a cat'
>>> # With sampling, the output becomes more creative!
>>> generated_ids = model.generate(**model_inputs, do_sample=True)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'I am a cat.\nI just need to be. I am always.\nEvery time'
์๋ชป๋ ํจ๋ฉ wrong-padding-side
LLM์ ๋์ฝ๋ ์ ์ฉ ๊ตฌ์กฐ๋ฅผ ๊ฐ์ง๊ณ ์์ด, ์
๋ ฅ ํ๋กฌํํธ์ ๋ํด ์ง์์ ์ผ๋ก ๋ฐ๋ณต ์ฒ๋ฆฌ๋ฅผ ํฉ๋๋ค. ์
๋ ฅ ๋ฐ์ดํฐ์ ๊ธธ์ด๊ฐ ๋ค๋ฅด๋ฉด ํจ๋ฉ ์์
์ด ํ์ํฉ๋๋ค. LLM์ ํจ๋ฉ ํ ํฐ์์ ์๋์ ์ด์ด๊ฐ๋๋ก ์ค๊ณ๋์ง ์์๊ธฐ ๋๋ฌธ์, ์
๋ ฅ ์ผ์ชฝ์ ํจ๋ฉ์ด ์ถ๊ฐ ๋์ด์ผ ํฉ๋๋ค. ๊ทธ๋ฆฌ๊ณ ์ดํ
์
๋ง์คํฌ๋ ๊ผญ generate ํจ์์ ์ ๋ฌ๋์ด์ผ ํฉ๋๋ค!
>>> # The tokenizer initialized above has right-padding active by default: the 1st sequence,
>>> # which is shorter, has padding on the right side. Generation fails.
>>> model_inputs = tokenizer(
... ["1, 2, 3", "A, B, C, D, E"], padding=True, return_tensors="pt"
... ).to("cuda")
>>> generated_ids = model.generate(**model_inputs)
>>> tokenizer.batch_decode(generated_ids[0], skip_special_tokens=True)[0]
''
>>> # With left-padding, it works as expected!
>>> tokenizer = AutoTokenizer.from_pretrained("openlm-research/open_llama_7b", padding_side="left")
>>> tokenizer.pad_token = tokenizer.eos_token # Llama has no pad token by default
>>> model_inputs = tokenizer(
... ["1, 2, 3", "A, B, C, D, E"], padding=True, return_tensors="pt"
... ).to("cuda")
>>> generated_ids = model.generate(**model_inputs)
>>> tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
'1, 2, 3, 4, 5, 6,'
์ถ๊ฐ ์๋ฃ further-resources
์๊ธฐํ๊ท ์์ฑ ํ๋ก์ธ์ค๋ ์๋์ ์ผ๋ก ๋จ์ํ ํธ์ด์ง๋ง, LLM์ ์ต๋ํ ํ์ฉํ๋ ค๋ฉด ์ฌ๋ฌ ๊ฐ์ง ์์๋ฅผ ๊ณ ๋ คํด์ผ ํ๋ฏ๋ก ์ฝ์ง ์์ ์ ์์ต๋๋ค. LLM์ ๋ํ ๋ ๊น์ ์ดํด์ ํ์ฉ์ ์ํ ๋ค์ ๋จ๊ณ๋ ์๋์ ๊ฐ์ต๋๋ค:
๊ณ ๊ธ ์์ฑ ์ฌ์ฉ advanced-generate-usage
- ๊ฐ์ด๋๋ ๋ค์ํ ์์ฑ ๋ฐฉ๋ฒ์ ์ ์ดํ๋ ๋ฐฉ๋ฒ, ์์ฑ ์ค์ ํ์ผ์ ์ค์ ํ๋ ๋ฐฉ๋ฒ, ์ถ๋ ฅ์ ์คํธ๋ฆฌ๋ฐํ๋ ๋ฐฉ๋ฒ์ ๋ํด ์ค๋ช ํฉ๋๋ค.
- [
~generation.GenerationConfig]์ [~generation.GenerationMixin.generate], generate-related classes๋ฅผ ์ฐธ์กฐํด๋ณด์ธ์.
LLM ๋ฆฌ๋๋ณด๋ llm-leaderboards
- Open LLM Leaderboard๋ ์คํ ์์ค ๋ชจ๋ธ์ ํ์ง์ ์ค์ ์ ๋ก๋๋ค.
- Open LLM-Perf Leaderboard๋ LLM ์ฒ๋ฆฌ๋์ ์ค์ ์ ๋ก๋๋ค.
์ง์ฐ ์๊ฐ ๋ฐ ์ฒ๋ฆฌ๋ latency-and-throughput
- ๋ฉ๋ชจ๋ฆฌ ์๊ตฌ ์ฌํญ์ ์ค์ด๋ ค๋ฉด, ๋์ ์์ํ์ ๋ํ ๊ฐ์ด๋๋ฅผ ์ฐธ์กฐํ์ธ์.
๊ด๋ จ ๋ผ์ด๋ธ๋ฌ๋ฆฌ related-libraries
text-generation-inference๋ LLM์ ์ํ ์ค์ ์ด์ ํ๊ฒฝ์ ์ ํฉํ ์๋ฒ์ ๋๋ค.optimum์ ํน์ ํ๋์จ์ด ์ฅ์น์์ LLM์ ์ต์ ํํ๊ธฐ ์ํด ๐ค Transformers๋ฅผ ํ์ฅํ ๊ฒ์ ๋๋ค.