Views
No views yet
1from transformers import AutoModel, MistralCommonBackend, Mistral3ForConditionalGeneration
2import torch
3
4# Load model and tokenizer
5model_id = "tiny-random/mistral-3"
6model = Mistral3ForConditionalGeneration.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)
13image_url = "https://static.wikia.nocookie.net/essentialsdocs/images/7/70/Battle.png/revision/latest?cb=20220523172438"
14messages = [
15 {
16 "role": "user",
17 "content": [
18 {
19 "type": "text",
20 "text": "What is this?",
21 },
22 {"type": "image_url", "image_url": {"url": image_url}},
23 ],
24 },
25]
26
27tokenized = tokenizer.apply_chat_template(
28 messages, return_tensors="pt", return_dict=True)
29tokenized["input_ids"] = tokenized["input_ids"].to(device="cuda")
30tokenized["pixel_values"] = tokenized["pixel_values"].to(
31 dtype=torch.bfloat16, device="cuda")
32image_sizes = [tokenized["pixel_values"].shape[-2:]]
33
34output = model.generate(
35 **tokenized.to("cuda"),
36 image_sizes=image_sizes,
37 max_new_tokens=32,
38)[0]
39
40decoded_output = tokenizer.decode(output[len(tokenized["input_ids"][0]):])
41print(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 set_seed,
13 Mistral3ForConditionalGeneration,
14 MistralCommonBackend,
15)
16
17source_model_id = "mistralai/Ministral-3-14B-Reasoning-2512"
18save_folder = "/tmp/tiny-random/mistral-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['text_config'].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})
37config_json['vision_config'].update({
38 "head_dim": 32,
39 "hidden_size": 128,
40 "intermediate_size": 128,
41 "num_attention_heads": 4,
42 "num_hidden_layers": 2,
43})
44with open(f"{save_folder}/config.json", "w", encoding='utf-8') as f:
45 json.dump(config_json, f, indent=2)
46
47config = AutoConfig.from_pretrained(
48 save_folder,
49 trust_remote_code=True,
50)
51print(config)
52torch.set_default_dtype(torch.bfloat16)
53model = Mistral3ForConditionalGeneration(config)
54torch.set_default_dtype(torch.float32)
55if file_exists(filename="generation_config.json", repo_id=source_model_id, repo_type='model'):
56 model.generation_config = GenerationConfig.from_pretrained(
57 source_model_id, trust_remote_code=True,
58 )
59 model.generation_config.do_sample = True
60 print(model.generation_config)
61model = model.cpu()
62with torch.no_grad():
63 for name, p in sorted(model.named_parameters()):
64 torch.nn.init.normal_(p, 0, 0.1)
65 print(name, p.shape)
66model.save_pretrained(save_folder)
67print(model)1Mistral3ForConditionalGeneration(
2 (model): Mistral3Model(
3 (vision_tower): PixtralVisionModel(
4 (patch_conv): Conv2d(3, 128, kernel_size=(14, 14), stride=(14, 14), bias=False)
5 (ln_pre): PixtralRMSNorm((128,), eps=1e-05)
6 (transformer): PixtralTransformer(
7 (layers): ModuleList(
8 (0-1): 2 x PixtralAttentionLayer(
9 (attention_norm): PixtralRMSNorm((128,), eps=1e-05)
10 (feed_forward): PixtralMLP(
11 (gate_proj): Linear(in_features=128, out_features=128, bias=False)
12 (up_proj): Linear(in_features=128, out_features=128, bias=False)
13 (down_proj): Linear(in_features=128, out_features=128, bias=False)
14 (act_fn): SiLUActivation()
15 )
16 (attention): PixtralAttention(
17 (k_proj): Linear(in_features=128, out_features=128, bias=False)
18 (v_proj): Linear(in_features=128, out_features=128, bias=False)
19 (q_proj): Linear(in_features=128, out_features=128, bias=False)
20 (o_proj): Linear(in_features=128, out_features=128, bias=False)
21 )
22 (ffn_norm): PixtralRMSNorm((128,), eps=1e-05)
23 )
24 )
25 )
26 (patch_positional_embedding): PixtralRotaryEmbedding()
27 )
28 (multi_modal_projector): Mistral3MultiModalProjector(
29 (norm): Mistral3RMSNorm((128,), eps=1e-05)
30 (patch_merger): Mistral3PatchMerger(
31 (merging_layer): Linear(in_features=512, out_features=128, bias=False)
32 )
33 (linear_1): Linear(in_features=128, out_features=8, bias=False)
34 (act): GELUActivation()
35 (linear_2): Linear(in_features=8, out_features=8, bias=False)
36 )
37 (language_model): Ministral3Model(
38 (embed_tokens): Embedding(131072, 8, padding_idx=11)
39 (layers): ModuleList(
40 (0-1): 2 x Ministral3DecoderLayer(
41 (self_attn): Ministral3Attention(
42 (q_proj): Linear(in_features=8, out_features=256, bias=False)
43 (k_proj): Linear(in_features=8, out_features=128, bias=False)
44 (v_proj): Linear(in_features=8, out_features=128, bias=False)
45 (o_proj): Linear(in_features=256, out_features=8, bias=False)
46 )
47 (mlp): Ministral3MLP(
48 (gate_proj): Linear(in_features=8, out_features=64, bias=False)
49 (up_proj): Linear(in_features=8, out_features=64, bias=False)
50 (down_proj): Linear(in_features=64, out_features=8, bias=False)
51 (act_fn): SiLUActivation()
52 )
53 (input_layernorm): Ministral3RMSNorm((8,), eps=1e-05)
54 (post_attention_layernorm): Ministral3RMSNorm((8,), eps=1e-05)
55 )
56 )
57 (norm): Ministral3RMSNorm((8,), eps=1e-05)
58 (rotary_emb): Ministral3RotaryEmbedding()
59 )
60 )
61 (lm_head): Linear(in_features=8, out_features=131072, bias=False)
62)