Views
No views yet
1import io
2
3import requests
4import torch
5from diffusers import Flux2Pipeline
6from diffusers.utils import load_image
7from huggingface_hub import get_token
8
9model_id = "tiny-random/flux.2"
10device = "cuda:0"
11torch_dtype = torch.bfloat16
12
13pipe = Flux2Pipeline.from_pretrained(
14 model_id, torch_dtype=torch_dtype
15).to(device)
16
17prompt = "Realistic macro photograph of a hermit crab using a soda can as its shell"
18cat_image = load_image(
19 "https://huggingface.co/spaces/zerogpu-aoti/FLUX.1-Kontext-Dev-fp8-dynamic/resolve/main/cat.png")
20image = pipe(
21 prompt=prompt,
22 image=[cat_image], # optional multi-image input
23 generator=torch.Generator(device=device).manual_seed(42),
24 num_inference_steps=4,
25 guidance_scale=4,
26 text_encoder_out_layers=(1,),
27).images[0]
28print(image)1import json
2
3import torch
4from diffusers import (
5 AutoencoderKLFlux2,
6 FlowMatchEulerDiscreteScheduler,
7 Flux2Pipeline,
8 Flux2Transformer2DModel,
9)
10from huggingface_hub import hf_hub_download
11from transformers import (
12 AutoConfig,
13 AutoTokenizer,
14 Mistral3ForConditionalGeneration,
15 PixtralProcessor,
16)
17from transformers.generation import GenerationConfig
18
19source_model_id = "black-forest-labs/FLUX.2-dev"
20save_folder = "/tmp/tiny-random/flux.2"
21
22torch.set_default_dtype(torch.bfloat16)
23scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
24 source_model_id, subfolder='scheduler')
25tokenizer = PixtralProcessor.from_pretrained(
26 source_model_id, subfolder='tokenizer')
27
28def save_json(path, obj):
29 import json
30 from pathlib import Path
31 Path(path).parent.mkdir(parents=True, exist_ok=True)
32 with open(path, 'w', encoding='utf-8') as f:
33 json.dump(obj, f, indent=2, ensure_ascii=False)
34
35def init_weights(model):
36 import torch
37 from transformers import set_seed
38 set_seed(42)
39 model = model.cpu()
40 with torch.no_grad():
41 for name, p in sorted(model.named_parameters()):
42 torch.nn.init.normal_(p, 0, 0.1)
43 print(name, p.shape, p.dtype, p.device)
44
45with open(hf_hub_download(source_model_id, filename='text_encoder/config.json', repo_type='model'), 'r', encoding='utf - 8') as f:
46 config = json.load(f)
47 config['text_config'].update({
48 'hidden_size': 8,
49 'intermediate_size': 64,
50 "head_dim": 32,
51 'num_attention_heads': 8,
52 'num_hidden_layers': 2,
53 'num_key_value_heads': 4,
54 'tie_word_embeddings': True,
55 })
56 config['vision_config'].update(
57 {
58 "head_dim": 32,
59 "hidden_size": 32,
60 "intermediate_size": 64,
61 "num_attention_heads": 1,
62 "num_hidden_layers": 2,
63 }
64 )
65 save_json(f'{save_folder}/text_encoder/config.json', config)
66 text_encoder_config = AutoConfig.from_pretrained(
67 f'{save_folder}/text_encoder')
68 text_encoder = Mistral3ForConditionalGeneration(
69 text_encoder_config).to(torch.bfloat16)
70 generation_config = GenerationConfig.from_pretrained(
71 source_model_id, subfolder='text_encoder')
72 # text_encoder.config.generation_config = generation_config
73 text_encoder.generation_config = generation_config
74 init_weights(text_encoder)
75
76with open(hf_hub_download(source_model_id, filename='transformer/config.json', repo_type='model'), 'r', encoding='utf-8') as f:
77 config = json.load(f)
78 config.update({
79 'attention_head_dim': 32,
80 "in_channels": 32,
81 'axes_dims_rope': [8, 12, 12],
82 'joint_attention_dim': 8,
83 'num_attention_heads': 2,
84 'num_layers': 2,
85 'num_single_layers': 2,
86 })
87 save_json(f'{save_folder}/transformer/config.json', config)
88 transformer_config = Flux2Transformer2DModel.load_config(
89 f'{save_folder}/transformer')
90 transformer = Flux2Transformer2DModel.from_config(transformer_config)
91 init_weights(transformer)
92
93with open(hf_hub_download(source_model_id, filename='vae/config.json', repo_type='model'), 'r', encoding='utf-8') as f:
94 config = json.load(f)
95 config.update({
96 'layers_per_block': 1,
97 'block_out_channels': [32, 32],
98 'latent_channels': 8,
99 'down_block_types': ['DownEncoderBlock2D', 'DownEncoderBlock2D'],
100 'up_block_types': ['UpDecoderBlock2D', 'UpDecoderBlock2D']
101 })
102 save_json(f'{save_folder}/vae/config.json', config)
103 vae_config = AutoencoderKLFlux2.load_config(f'{save_folder}/vae')
104 vae = AutoencoderKLFlux2.from_config(vae_config)
105 init_weights(vae)
106
107pipeline = Flux2Pipeline(
108 scheduler=scheduler,
109 text_encoder=text_encoder,
110 tokenizer=tokenizer,
111 transformer=transformer,
112 vae=vae,
113)
114pipeline = pipeline.to(torch.bfloat16)
115pipeline.save_pretrained(save_folder, safe_serialization=True)
116print(pipeline)