Views
No views yet

user_embedding.safetensors contains embeddings for 1000 training users (BF16, shape [1000, 30720]). The users/ and users_linear/ directories contain individually fine-tuned weights for 50 test users (IDs 3685~4279), which are distinct from the training set.| File | Description | Size |
|---|---|---|
mod_adapter.safetensors | Modulation adapter weights (trained at 260k steps) | ≈1.81 GB |
user_embedding.safetensors | Shared user preference embedding (1000 training users) | ≈61 MB |
adapter_config.yaml | Adapter architecture configuration | - |
users/user_embedding_*.safetensors | Per-user embedding weights — 50 test users (non-linear) | ≈60 KB each |
users_linear/user_combination_*.safetensors | Linear combination user weights — 50 test users (linear) | ≈2 KB each |
1git clone https://github.com/120L020904/Premier.git
2cd Premier
3pip install -r requirements.txt1# Download from HuggingFace (assume saved to ./Premier/)
2from huggingface_hub import snapshot_download
3snapshot_download("pino10010/Premier", local_dir="./Premier")./Premier/
├── adapter_config.yaml
├── mod_adapter.safetensors
├── user_embedding.safetensors
├── users/
│ └── user_embedding_*.safetensors
└── users_linear/
└── user_combination_*.safetensors1import os
2import sys
3import torch
4from diffusers import FluxPipeline
5from safetensors.torch import load_file
6from torch import nn
7
8sys.path.append("path/to/Premier")
9from scripts.pipeline.flux_adapter import generate_xverse
10from scripts.pipeline.mod_adapters import load_modulation_adapter
11from scripts.utils.utils import get_config, save_images
12
13device = "cuda"
14dtype = torch.bfloat16
15model_dir = "./Premier"
16
17# Load FLUX.1-dev base model
18pipe = FluxPipeline.from_pretrained(
19 "black-forest-labs/FLUX.1-dev",
20 torch_dtype=dtype
21).to(device)
22
23# Load adapter config and weights
24adapter_config = get_config(config_path=os.path.join(model_dir, "adapter_config.yaml"))
25mod_adapter = load_modulation_adapter(
26 adapter_config, dtype, device,
27 ckpt_dir=model_dir,
28 is_training=False
29)
30mod_adapter.eval()
31
32# Load shared user embedding (1000 training users, BF16, shape [1000, 30720])
33user_token_num = adapter_config["model"]["modulation"]["user_token_num"] # 30
34state_dict = load_file(os.path.join(model_dir, "user_embedding.safetensors"))
35user_embedding = nn.Embedding(
36 num_embeddings=1000,
37 embedding_dim=user_token_num * 1024
38).to(device=device, dtype=dtype)
39user_embedding.load_state_dict(state_dict)
40
41# Generate image for a training-set user (IDs 0~999)
42user_id = 0
43indices = torch.tensor([user_id], dtype=torch.long).to(device)
44user_pref = user_embedding(indices).view(-1, user_token_num, 1024)
45
46prompt = "a cute cat sitting on a windowsill in watercolor style"
47generator = torch.Generator(device).manual_seed(42)
48
49result = generate_xverse(
50 pipeline=pipe,
51 mod_adapter=mod_adapter,
52 user_preference_embedding=user_pref,
53 prompt=prompt,
54 prompt_2=prompt,
55 num_inference_steps=30,
56 guidance_scale=2.5,
57 height=512,
58 width=512,
59 generator=generator,
60 model_config=adapter_config,
61)
62image = result.images[0]
63image.save("output.png")1# Load individual user embedding (non-linear, for user 3685)
2user_id = 3685
3user_weights = load_file(os.path.join(model_dir, f"users/user_embedding_{user_id}.safetensors"))
4user_token_num = 30
5
6train_user_embedding = nn.Embedding(
7 num_embeddings=1,
8 embedding_dim=user_token_num * 1024
9).to(device=device, dtype=dtype)
10train_user_embedding.load_state_dict(user_weights)
11
12indices = torch.tensor([0], dtype=torch.long).to(device)
13user_pref = train_user_embedding(indices).view(-1, user_token_num, 1024)
14
15# Use same generate_xverse() call as above with user_pref1from scripts.train_flux.train_user_embedding_linear import EmbeddingLinearCombination
2
3user_id = 3685
4user_token_num = 30
5
6# Load shared training embedding
7train_state_dict = load_file(os.path.join(model_dir, "user_embedding.safetensors"))
8train_user_embedding = nn.Embedding(
9 num_embeddings=1000,
10 embedding_dim=user_token_num * 1024
11).to(device=device, dtype=dtype)
12train_user_embedding.load_state_dict(train_state_dict)
13
14# Load linear combination weights
15combination_state_dict = load_file(
16 os.path.join(model_dir, f"users_linear/user_combination_{user_id}.safetensors")
17)
18embedding_comb = EmbeddingLinearCombination(
19 combination_size=1, embedding_num=1000, use_softmax=False
20).to(device=device, dtype=dtype)
21embedding_comb.load_state_dict(combination_state_dict)
22
23indices = torch.tensor([0], dtype=torch.long).to(device)
24user_pref = embedding_comb(
25 train_user_embedding, input_ids=indices
26).view(-1, user_token_num, 1024)
27
28# Use same generate_xverse() call as above with user_pref1@article{wang2026premier,
2 title={Premier: Personalized Preference Modulation with Learnable User Embedding in Text-to-Image Generation},
3 author={Wang, Zihao and Wei, Yuxiang and Zhou, Xinpeng and Zhang, Tianyu and Liang, Tao and Bai, Yalong and Zhang, Hongzhi and Zuo, Wangmeng},
4 journal={arXiv preprint arXiv:2603.20725},
5 year={2026}
6}