Views
No views yet
1# init parameters
2model_name: str = 'scb10x/typhoon-7b'
3quantization_mode: str = 'q4-bnb_cuda' # possible values = {'q4-bnb_cuda', 'q8-bnb_cuda', 'q4-torch_ptdq', 'q8-torch_ptdq'}
4
5# load tokenizer
6from transformers import AutoTokenizer
7
8tokenizer = AutoTokenizer.from_pretrained(model_name)
9tokenizer.pad_token_id = tokenizer.eos_token_id
10print(tokenizer) # LlamaTokenizerFast
11
12# load model
13import torch
14from transformers import AutoModelForCausalLM
15
16if quantization_mode == 'q4-bnb_cuda': # ampere architecture with 8gb vram + cpu with 20gb is recommended
17 print('4-bits bitsandbytes quantization with cuda')
18 model = AutoModelForCausalLM.from_pretrained(
19 model_name,
20 load_in_4bit = True,
21 device_map = 'auto',
22 torch_dtype = torch.bfloat16)
23elif quantization_mode == 'q8-bnb_cuda': # ampere architecture with 12gb vram + cpu with 20gb is recommended
24 print('8-bits bitsandbytes quantization with cuda')
25 model = AutoModelForCausalLM.from_pretrained(
26 model_name,
27 load_in_8bit = True,
28 device_map = 'auto',
29 torch_dtype = torch.bfloat16)
30elif quantization_mode == 'q4-torch_ptdq': # cpu with 64gb++ ram is recommended
31 print('4-bits x2 post training dynamic quantization')
32 base_model = AutoModelForCausalLM.from_pretrained(
33 model_name,
34 torch_dtype = torch.float32)
35 model = torch.quantization.quantize_dynamic(base_model, dtype = torch.quint4x2)
36elif quantization_mode == 'q8-torch_ptdq': # cpu with 64gb++ ram is recommended
37 print('8-bits post training dynamic quantization')
38 base_model = AutoModelForCausalLM.from_pretrained(
39 model_name,
40 torch_dtype = torch.float32)
41 model = torch.quantization.quantize_dynamic(base_model, dtype = torch.quint8)
42else:
43 print('default model')
44 model = AutoModelForCausalLM.from_pretrained(model_name)
45print(model) # MistralForCausalLM
46
47# text generator
48from transformers import GenerationConfig, TextGenerationPipeline
49
50config = GenerationConfig.from_pretrained(model_name)
51config.num_return_sequences: int = 1
52config.do_sample: bool = True
53config.max_new_tokens: int = 128
54config.temperature: float = 0.7
55config.top_p: float = 0.95
56config.repetition_penalty: float = 1.3
57generator = TextGenerationPipeline(
58 model = model,
59 tokenizer = tokenizer,
60 return_full_text = True,
61 generation_config = config)
62
63# sample
64sample: str = 'ความหมายของชีวิตคืออะไร?\n'
65output = generator(sample, pad_token_id = tokenizer.eos_token_id)
66print(output[0]['generated_text'])requirement.txt1torch==2.1.2
2accelerate==0.25.0
3bitsandbytes==0.41.3
4#transformers==4.37.0.dev0
5transformers @ git+https://github.com/huggingface/transformers