Views
No views yet
1import torch
2from transformers import Ministral3ForCausalLM, MistralCommonBackend
3
4# Load model and tokenizer
5model_id = "tiny-random/ministral-3"
6model = Ministral3ForCausalLM.from_pretrained(
7 model_id,
8 device_map="cuda",
9 torch_dtype="bfloat16",
10 trust_remote_code=True,
11)
12tokenizer = MistralCommonBackend.from_pretrained(model_id)
13messages = [
14 {
15 "role": "user",
16 "content": "Hi",
17 },
18]
19
20tokenized = tokenizer.apply_chat_template(
21 messages, return_tensors="pt", return_dict=True)
22output = model.generate(
23 **tokenized.to("cuda"),
24 max_new_tokens=32,
25)[0]
26decoded_output = tokenizer.decode(output[len(tokenized["input_ids"][0]):])
27print(decoded_output)1import json
2from pathlib import Path
3
4import accelerate
5import torch
6from huggingface_hub import file_exists, hf_hub_download
7from transformers import (
8 AutoConfig,
9 AutoModelForCausalLM,
10 AutoProcessor,
11 GenerationConfig,
12 Ministral3ForCausalLM,
13 MistralCommonBackend,
14 set_seed,
15)
16
17source_model_id = "mistralai/Devstral-2-123B-Instruct-2512"
18save_folder = "/tmp/tiny-random/ministral-3"
19
20processor = AutoProcessor.from_pretrained(
21 source_model_id, trust_remote_code=True)
22processor.save_pretrained(save_folder)
23processor = MistralCommonBackend.from_pretrained(
24 source_model_id, trust_remote_code=True)
25processor.save_pretrained(save_folder)
26
27with open(hf_hub_download(source_model_id, filename='config.json', repo_type='model'), 'r', encoding='utf-8') as f:
28 config_json = json.load(f)
29config_json.update({
30 "head_dim": 32,
31 "hidden_size": 8,
32 "intermediate_size": 64,
33 "num_attention_heads": 8,
34 "num_hidden_layers": 2,
35 "num_key_value_heads": 4,
36 "tie_word_embeddings": True,
37})
38del config_json['quantization_config']
39with open(f"{save_folder}/config.json", "w", encoding='utf-8') as f:
40 json.dump(config_json, f, indent=2)
41
42config = AutoConfig.from_pretrained(
43 save_folder,
44 trust_remote_code=True,
45)
46print(config)
47torch.set_default_dtype(torch.bfloat16)
48model = Ministral3ForCausalLM(config)
49torch.set_default_dtype(torch.float32)
50if file_exists(filename="generation_config.json", repo_id=source_model_id, repo_type='model'):
51 model.generation_config = GenerationConfig.from_pretrained(
52 source_model_id, trust_remote_code=True,
53 )
54 model.generation_config.do_sample = True
55 print(model.generation_config)
56model = model.cpu()
57with torch.no_grad():
58 for name, p in sorted(model.named_parameters()):
59 torch.nn.init.normal_(p, 0, 0.1)
60 print(name, p.shape)
61model.save_pretrained(save_folder)
62print(model)1Ministral3ForCausalLM(
2 (model): Ministral3Model(
3 (embed_tokens): Embedding(131072, 8, padding_idx=11)
4 (layers): ModuleList(
5 (0-1): 2 x Ministral3DecoderLayer(
6 (self_attn): Ministral3Attention(
7 (q_proj): Linear(in_features=8, out_features=256, bias=False)
8 (k_proj): Linear(in_features=8, out_features=128, bias=False)
9 (v_proj): Linear(in_features=8, out_features=128, bias=False)
10 (o_proj): Linear(in_features=256, out_features=8, bias=False)
11 )
12 (mlp): Ministral3MLP(
13 (gate_proj): Linear(in_features=8, out_features=64, bias=False)
14 (up_proj): Linear(in_features=8, out_features=64, bias=False)
15 (down_proj): Linear(in_features=64, out_features=8, bias=False)
16 (act_fn): SiLUActivation()
17 )
18 (input_layernorm): Ministral3RMSNorm((8,), eps=1e-05)
19 (post_attention_layernorm): Ministral3RMSNorm((8,), eps=1e-05)
20 )
21 )
22 (norm): Ministral3RMSNorm((8,), eps=1e-05)
23 (rotary_emb): Ministral3RotaryEmbedding()
24 )
25 (lm_head): Linear(in_features=8, out_features=131072, bias=False)
26)