Coverage for wrapper/januspro/wrapper_januspro.py: 91%
183 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1"""
2https://github.com/deepseek-ai/Janus/blob/1daa72fa409002d40931bd7b36a9280362469ead/demo/app_januspro.py#L15
3"""
4import logging
5import os
6import sys
7import random
9from typing import override
10from typing import Optional
11from typing import Dict
12from typing import Any
13from typing import Union
15from PIL import Image
17import numpy as np
19import torch
20import torch.distributed as dist
21from torch import inference_mode
23from wrapper_model import ModelGeneration
25from transformers import AutoModelForCausalLM
26from transformers import AutoConfig
27from janus.models import VLChatProcessor
29from xfuser.config import EngineConfig
32class JanusProGeneration(ModelGeneration):
33 """Wrapper class for Janus Pro model generation."""
35 def __init__(
36 self,
37 model_name: str = "januspro",
38 engine_config: EngineConfig = None,
39 param_dtype: torch.dtype = torch.bfloat16,
40 ) -> None:
41 super().__init__(model_name)
43 self.engine_config = engine_config
44 if self.engine_config is not None:
45 self.torch_compile = self.engine_config.runtime_config.use_torch_compile
46 self.param_dtype = param_dtype
48 # Parallelism
49 self.gpu: Optional[str] = None
50 if torch.cuda.is_available():
51 self.gpu = torch.cuda.get_device_name(0)
53 self.base_seed = random.randint(0, sys.maxsize)
55 # Model components
56 self.vl_gpt: Optional[torch.nn.Module] = None
57 self.vl_chat_processor: Optional[Any] = None
58 self.tokenizer: Optional[Any] = None
60 def __del__(self) -> None:
61 # Clean models
62 if self.vl_gpt is not None:
63 self.vl_gpt = None
64 if self.vl_chat_processor is not None:
65 self.vl_chat_processor = None
66 if self.tokenizer is not None:
67 self.tokenizer = None
68 if dist.is_initialized():
69 dist.destroy_process_group()
71 def init_parallelism(self) -> None:
72 self.load_timer.start("torch_dist")
74 self.rank = int(os.getenv("RANK", 0))
75 self.local_rank = int(os.getenv("LOCAL_RANK", 0))
76 self.world_size = int(os.getenv("WORLD_SIZE", 1))
78 self.device_id = self.local_rank
79 self.device = torch.device(f"cuda:{self.device_id}")
81 torch.cuda.set_device(self.local_rank)
83 if self.world_size > 1:
84 logging.warning("Janus is not optimized for multi-GPU setups (yet).")
85 self.world_size = 1
87 self.load_timer.end("torch_dist")
89 def load_model(self) -> None:
90 assert torch.cuda.is_available()
92 self.load_timer.start("processor")
93 self.MODEL_NAME = "deepseek-ai/Janus-Pro-7B"
94 self.vl_chat_processor = VLChatProcessor.from_pretrained(
95 self.MODEL_NAME
96 )
97 self.tokenizer = self.vl_chat_processor.tokenizer
98 self.load_timer.end("processor")
100 self.load_timer.start("config")
101 config = AutoConfig.from_pretrained(self.MODEL_NAME) # nosec B615
102 language_config = config.language_config
103 language_config._attn_implementation = 'eager'
104 self.load_timer.end("config")
106 self.load_timer.start("model")
107 self.vl_gpt = AutoModelForCausalLM.from_pretrained(
108 self.MODEL_NAME,
109 language_config=language_config,
110 trust_remote_code=True
111 ) # nosec B615
112 assert self.vl_gpt is not None
113 self.vl_gpt = self.vl_gpt.to(self.param_dtype) # type: ignore[arg-type]
114 self.vl_gpt = self.vl_gpt.to(self.device)
115 self.vl_gpt = self.vl_gpt.eval()
116 self.load_timer.end("model")
118 logging.info(f"Loaded Janus Pro: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.")
120 def init_model_parallelism(self) -> None:
121 if self.world_size > 1:
122 logging.warning("Janus Pro does not support model parallelism yet.")
124 def model_compile(self) -> None:
125 if not self.torch_compile:
126 return
128 self.load_timer.start("model_compile")
129 torch._inductor.config.reorder_for_compute_comm_overlap = True
130 # Note: Janus has complex architecture, be careful with compilation
131 # self.vl_gpt = torch.compile(self.vl_gpt, mode="max-autotune-no-cudagraphs")
132 self.load_timer.end("model_compile")
134 def _assert_model_init(self) -> None:
135 super()._assert_model_init()
136 assert self.vl_gpt is not None
137 assert self.vl_chat_processor is not None
138 assert self.tokenizer is not None
140 def _assert_args(
141 self,
142 img_size: int,
143 patch_size: int,
144 ) -> None:
145 if img_size % patch_size != 0:
146 raise ValueError(f"Image size {img_size} must be divisible by patch size {patch_size}")
147 if img_size < 384:
148 raise ValueError(f"Image size {img_size} must be at least 384")
150 def _prepare_prompt(self, prompt: str) -> str:
151 assert self.vl_chat_processor is not None
152 messages = [
153 {'role': '<|User|>', 'content': prompt},
154 {'role': '<|Assistant|>', 'content': ''}
155 ]
156 text = self.vl_chat_processor.apply_sft_template_for_multi_turn_prompts(
157 conversations=messages,
158 sft_format=self.vl_chat_processor.sft_format,
159 system_prompt=''
160 )
161 return text + self.vl_chat_processor.image_start_tag
163 @inference_mode()
164 async def warmup(self) -> None:
165 logging.info(f"[{self.rank}] Warmup for Janus Pro generation.")
166 await self.generate(
167 prompt="A warmup image to initialize the model.",
168 img_size=384,
169 image_token_num_per_image=576
170 )
172 @override
173 @inference_mode()
174 async def generate(
175 self,
176 prompt: str,
177 temperature: float = 1.0,
178 cfg_weight: float = 5.0,
179 image_token_num_per_image: int = 576,
180 img_size: int = 384,
181 patch_size: int = 16,
182 job_id: Optional[str] = None,
183 ) -> Image.Image:
184 """
185 Generate images from a prompt using the Janus Pro model.
186 Args:
187 prompt (str): Text prompt to guide the image generation.
188 temperature (float): Sampling temperature for generation.
189 parallel_size (int): Number of images to generate in parallel.
190 cfg_weight (float): Classifier-free guidance weight.
191 image_token_num_per_image (int): Number of tokens per image.
192 img_size (int): Size of the generated images.
193 patch_size (int): Patch size for the vision model.
194 Returns:
195 list[Image.Image]: List of generated PIL Images.
196 """
197 gen_timer = self._new_gen_timer(job_id)
199 self._assert_model_init()
200 assert self.vl_gpt is not None
201 assert self.vl_chat_processor is not None
202 assert self.tokenizer is not None
203 self._assert_args(img_size, patch_size)
205 width = img_size // patch_size * patch_size
206 height = img_size // patch_size * patch_size
208 # Single image generation for now
209 parallel_size = 1
211 self.running = True # Mark running to avoid concurrent calls
213 try:
214 torch.cuda.empty_cache()
216 gen_timer.start("prepare_prompt")
217 formatted_prompt = self._prepare_prompt(prompt)
218 gen_timer.end("prepare_prompt")
220 gen_timer.start("tokenize")
221 input_ids = torch.LongTensor(self.tokenizer.encode(formatted_prompt))
222 tokens = torch.zeros((parallel_size * 2, len(input_ids)), dtype=torch.int).to(self.device)
223 for i in range(parallel_size * 2):
224 tokens[i, :] = input_ids
225 if i % 2 != 0:
226 tokens[i, 1:-1] = self.vl_chat_processor.pad_id
227 get_input_emb = self.vl_gpt.language_model.get_input_embeddings # type: ignore[union-attr]
228 inputs_embeds = get_input_emb()(tokens) # type: ignore[operator]
229 gen_timer.end("tokenize")
231 gen_timer.start("generate_tokens")
232 generated_tokens = torch.zeros((parallel_size, image_token_num_per_image), dtype=torch.int).to(self.device)
233 pkv = None
234 for ix in range(image_token_num_per_image):
235 gen_timer.start(f"generate_token_{ix:03d}")
236 outputs = self.vl_gpt.language_model.model( # type: ignore[operator, union-attr]
237 inputs_embeds=inputs_embeds,
238 use_cache=True,
239 past_key_values=pkv
240 )
241 pkv = outputs.past_key_values
242 hidden_states = outputs.last_hidden_state
243 logits = self.vl_gpt.gen_head(hidden_states[:, -1, :]) # type: ignore[operator]
244 logit_cond = logits[0::2, :]
245 logit_uncond = logits[1::2, :]
246 logits = logit_uncond + cfg_weight * (logit_cond - logit_uncond)
247 probs = torch.softmax(logits / temperature, dim=-1)
248 next_token = torch.multinomial(probs, num_samples=1)
249 generated_tokens[:, ix] = next_token.squeeze(dim=-1)
250 next_token = torch.cat([
251 next_token.unsqueeze(dim=1),
252 next_token.unsqueeze(dim=1)
253 ], dim=1).view(-1)
255 img_embeds = self.vl_gpt.prepare_gen_img_embeds(next_token) # type: ignore[operator]
256 inputs_embeds = img_embeds.unsqueeze(dim=1)
257 gen_timer.end(f"generate_token_{ix:03d}")
258 gen_timer.end("generate_tokens")
260 gen_timer.start("decode_images")
261 # TODO fix failure
262 # shape '[2, 24, 24, 8]' is invalid for input of size 1600.
263 patches = self.vl_gpt.gen_vision_model.decode_code( # type: ignore[operator, union-attr]
264 generated_tokens.to(dtype=torch.int),
265 shape=[parallel_size, 8, width // patch_size, height // patch_size]
266 )
267 dec = patches.to(torch.float32).cpu().numpy().transpose(0, 2, 3, 1)
268 dec = np.clip((dec + 1) / 2 * 255, 0, 255)
269 visual_img = np.zeros((parallel_size, width, height, 3), dtype=np.uint8)
270 visual_img[:, :, :] = dec
271 gen_timer.end("decode_images")
273 gen_timer.start("convert_pil")
274 images = []
275 for i in range(parallel_size):
276 pil_image = Image.fromarray(visual_img[i]).resize((768, 768), Image.Resampling.LANCZOS)
277 images.append(pil_image)
278 gen_timer.end("convert_pil")
280 logging.info(f"[{self.rank}] Generated {len(images)} images. Return just 1.")
282 return images[0]
283 finally:
284 self.running = False
285 gen_timer.end("total")
287 def get_health(self) -> Dict[str, Any]:
288 ret = super().get_health()
289 ret.update({
290 "gpu": self.gpu,
291 "rank": self.rank,
292 "world_size": self.world_size,
293 "torch_compile": self.torch_compile,
294 "dtype": str(self.param_dtype),
295 })
296 return ret
298 async def get_rest_args(
299 self,
300 data_json: Dict[str, Union[str, int, float]]
301 ) -> Dict[str, Any]:
302 if data_json is None or not isinstance(data_json, dict):
303 raise ValueError("Missing JSON body")
305 prompt = data_json.get("prompt", None)
306 if prompt is None:
307 raise ValueError("Missing 'prompt' parameter")
309 temperature = float(data_json.get("temperature", 1.0))
310 cfg_weight = float(data_json.get("cfg_weight", 5.0))
311 image_token_num_per_image = int(data_json.get("image_token_num_per_image", 576))
312 img_size = int(data_json.get("img_size", 384))
313 patch_size = int(data_json.get("patch_size", 16))
315 return {
316 "task": self.model_name,
317 "args": {
318 "prompt": prompt,
319 "temperature": temperature,
320 "cfg_weight": cfg_weight,
321 "image_token_num_per_image": image_token_num_per_image,
322 "img_size": img_size,
323 "patch_size": patch_size,
324 }
325 }