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

1""" 

2https://github.com/deepseek-ai/Janus/blob/1daa72fa409002d40931bd7b36a9280362469ead/demo/app_januspro.py#L15 

3""" 

4import logging 

5import os 

6import sys 

7import random 

8 

9from typing import override 

10from typing import Optional 

11from typing import Dict 

12from typing import Any 

13from typing import Union 

14 

15from PIL import Image 

16 

17import numpy as np 

18 

19import torch 

20import torch.distributed as dist 

21from torch import inference_mode 

22 

23from wrapper_model import ModelGeneration 

24 

25from transformers import AutoModelForCausalLM 

26from transformers import AutoConfig 

27from janus.models import VLChatProcessor 

28 

29from xfuser.config import EngineConfig 

30 

31 

32class JanusProGeneration(ModelGeneration): 

33 """Wrapper class for Janus Pro model generation.""" 

34 

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) 

42 

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 

47 

48 # Parallelism 

49 self.gpu: Optional[str] = None 

50 if torch.cuda.is_available(): 

51 self.gpu = torch.cuda.get_device_name(0) 

52 

53 self.base_seed = random.randint(0, sys.maxsize) 

54 

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 

59 

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() 

70 

71 def init_parallelism(self) -> None: 

72 self.load_timer.start("torch_dist") 

73 

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)) 

77 

78 self.device_id = self.local_rank 

79 self.device = torch.device(f"cuda:{self.device_id}") 

80 

81 torch.cuda.set_device(self.local_rank) 

82 

83 if self.world_size > 1: 

84 logging.warning("Janus is not optimized for multi-GPU setups (yet).") 

85 self.world_size = 1 

86 

87 self.load_timer.end("torch_dist") 

88 

89 def load_model(self) -> None: 

90 assert torch.cuda.is_available() 

91 

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") 

99 

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") 

105 

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") 

117 

118 logging.info(f"Loaded Janus Pro: {self.MODEL_NAME} device:{self.device} dtype:{self.param_dtype}.") 

119 

120 def init_model_parallelism(self) -> None: 

121 if self.world_size > 1: 

122 logging.warning("Janus Pro does not support model parallelism yet.") 

123 

124 def model_compile(self) -> None: 

125 if not self.torch_compile: 

126 return 

127 

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") 

133 

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 

139 

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") 

149 

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 

162 

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 ) 

171 

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) 

198 

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) 

204 

205 width = img_size // patch_size * patch_size 

206 height = img_size // patch_size * patch_size 

207 

208 # Single image generation for now 

209 parallel_size = 1 

210 

211 self.running = True # Mark running to avoid concurrent calls 

212 

213 try: 

214 torch.cuda.empty_cache() 

215 

216 gen_timer.start("prepare_prompt") 

217 formatted_prompt = self._prepare_prompt(prompt) 

218 gen_timer.end("prepare_prompt") 

219 

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") 

230 

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) 

254 

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") 

259 

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") 

272 

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") 

279 

280 logging.info(f"[{self.rank}] Generated {len(images)} images. Return just 1.") 

281 

282 return images[0] 

283 finally: 

284 self.running = False 

285 gen_timer.end("total") 

286 

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 

297 

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") 

304 

305 prompt = data_json.get("prompt", None) 

306 if prompt is None: 

307 raise ValueError("Missing 'prompt' parameter") 

308 

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)) 

314 

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 }