Views
No views yet
DiNa-LRM: A diffusion-native latent reward model. It achieves competitive reward accuracy while being significantly cheaper for alignment by operating directly in the latent space.
diffusion_rm package installed:1# 1. Create a new conda environment
2conda create -n diffusion-rm python=3.10 -y
3conda activate diffusion-rm
4
5# 2. Install the package in editable mode
6# This will install all necessary dependencies including torch and diffusers
7pip install git+https://github.com/HKUST-C4G/diffusion-rm.git1import torch
2from diffusers import StableDiffusion3Pipeline
3from diffusion_rm.models.sd3_rm import encode_prompt
4from diffusion_rm.infer.inference import DRMInferencer
5
6# Load SD3.5 Pipeline
7device = torch.device('cuda:0')
8dtype = torch.bfloat16
9pipe = StableDiffusion3Pipeline.from_pretrained(
10 "stabilityai/stable-diffusion-3.5-medium",
11 torch_dtype=dtype
12).to(device)
13pipe.vae.to(device, dtype=dtype)
14pipe.text_encoder.to(device, dtype=dtype)
15pipe.text_encoder_2.to(device, dtype=dtype)
16pipe.text_encoder_3.to(device, dtype=dtype)
17pipe.transformer.to(device, dtype=dtype)
18
19text_encoders = [pipe.text_encoder, pipe.text_encoder_2, pipe.text_encoder_3]
20tokenizers = [pipe.tokenizer, pipe.tokenizer_2, pipe.tokenizer_3]
21
22def compute_text_embeddings(text_encoders, tokenizers, prompts):
23 with torch.no_grad():
24 prompt_embeds, pooled_prompt_embeds = encode_prompt(
25 text_encoders, tokenizers, prompts, max_sequence_length=256
26 )
27 prompt_embeds = prompt_embeds.to(text_encoders[0].device)
28 pooled_prompt_embeds = pooled_prompt_embeds.to(text_encoders[0].device)
29 return prompt_embeds, pooled_prompt_embeds
30
31# Initialize DiNa-LRM Scorer
32scorer = DRMInferencer(
33 pipeline=pipe,
34 config_path=None,
35 model_path="liuhuohuo/DiNa-LRM-SD35M-12layers",
36 device=device,
37 model_dtype=dtype,
38 load_from_disk=False,
39)
40
41# 1. Generate latents (Set output_type='latent' for DiNa-LRM)
42prompt = "A girl walking in the street"
43
44with torch.no_grad():
45 # Helper to get embeddings
46 prompt_embeds, pooled_embeds = compute_text_embeddings(text_encoders, tokenizers, [prompt])
47
48 output = pipe(
49 prompt_embeds=prompt_embeds,
50 pooled_prompt_embeds=pooled_embeds,
51 num_inference_steps=40,
52 guidance_scale=4.5,
53 output_type='latent'
54 )
55 latents = output.images
56
57
58# 2. Compute reward
59with torch.no_grad():
60 raw_score = scorer.reward(
61 text_conds={'encoder_hidden_states': prompt_embeds, 'pooled_projections': pooled_embeds},
62 latents=latents,
63 u=0.4
64 )
65 score = (raw_score + 10.0) / 10.0
66 print(f"DiNa-LRM Score: {score.item()}")
67
68# 3. [Optional] decode and save images
69with torch.no_grad():
70 latents_decoded = (latents / pipe.vae.config.scaling_factor) + pipe.vae.config.shift_factor
71 image = pipe.vae.decode(latents_decoded.to(pipe.vae.dtype), return_dict=False)[0]
72 image = pipe.image_processor.postprocess(image, output_type="pil")[0]
73
74image.save("example.png")1import torch
2import torchvision.transforms as T
3from PIL import Image
4from diffusers import StableDiffusion3Pipeline
5from diffusion_rm.models.sd3_rm import encode_prompt
6from diffusion_rm.infer.inference import DRMInferencer
7
8# Load SD3.5 Pipeline
9device = torch.device('cuda:0')
10dtype = torch.bfloat16
11pipe = StableDiffusion3Pipeline.from_pretrained(
12 "stabilityai/stable-diffusion-3.5-medium",
13 torch_dtype=dtype
14).to(device)
15pipe.vae.to(device, dtype=dtype)
16pipe.text_encoder.to(device, dtype=dtype)
17pipe.text_encoder_2.to(device, dtype=dtype)
18pipe.text_encoder_3.to(device, dtype=dtype)
19pipe.transformer.to(device, dtype=dtype)
20
21text_encoders = [pipe.text_encoder, pipe.text_encoder_2, pipe.text_encoder_3]
22tokenizers = [pipe.tokenizer, pipe.tokenizer_2, pipe.tokenizer_3]
23
24def compute_text_embeddings(text_encoders, tokenizers, prompts):
25 with torch.no_grad():
26 prompt_embeds, pooled_prompt_embeds = encode_prompt(
27 text_encoders, tokenizers, prompts, max_sequence_length=256
28 )
29 prompt_embeds = prompt_embeds.to(text_encoders[0].device)
30 pooled_prompt_embeds = pooled_prompt_embeds.to(text_encoders[0].device)
31 return prompt_embeds, pooled_prompt_embeds
32
33# Initialize DiNa-LRM Scorer
34scorer = DRMInferencer(
35 pipeline=pipe,
36 config_path=None,
37 model_path="liuhuohuo/DiNa-LRM-SD35M-12layers",
38 device=device,
39 model_dtype=dtype,
40 load_from_disk=False,
41)
42
43# 1. Load and Preprocess Image
44image_path = "assets/example.png"
45raw_image = Image.open(image_path).convert("RGB")
46transform = T.Compose([
47 T.ToTensor(),
48 T.Normalize([0.5], [0.5])
49])
50image_tensor = transform(raw_image).unsqueeze(0).to(device, dtype=dtype)
51
52prompt = "A girl walking in the street"
53
54with torch.no_grad():
55 # Helper to get embeddings
56 prompt_embeds, pooled_embeds = compute_text_embeddings(text_encoders, tokenizers, [prompt])
57
58
59# 2. Encode to Latent Space
60with torch.no_grad():
61 latents = pipe.vae.encode(image_tensor).latent_dist.sample()
62 # Apply SD3-specific scaling and shift
63 latents = (latents - pipe.vae.config.shift_factor) * pipe.vae.config.scaling_factor
64
65# 3. Compute Reward
66# Note: score normalization is often calculated as: score = (raw_score + 10.0) / 10.0
67raw_score = scorer.reward(
68 text_conds={'encoder_hidden_states': prompt_embeds, 'pooled_projections': pooled_embeds},
69 latents=latents,
70 u=0.1 # Lower u is recommended for static/clean images
71)
72score = (raw_score + 10.0) / 10.0
73print(f"Local Image Score: {score.item()}")
74