(args)
| 478 | |
| 479 | |
| 480 | def build_llava_engine(args): |
| 481 | processor = AutoProcessor.from_pretrained(args.model_path) |
| 482 | if args.model_type == "llava": |
| 483 | raw_image = Image.new('RGB', [10, 10]) # dummy image |
| 484 | image = processor(text="dummy", images=raw_image, |
| 485 | return_tensors="pt")['pixel_values'].to( |
| 486 | args.device, torch.float16) |
| 487 | |
| 488 | class LlavaVisionWrapper(torch.nn.Module): |
| 489 | |
| 490 | def __init__(self, tower, projector, feature_layer): |
| 491 | super().__init__() |
| 492 | self.tower = tower |
| 493 | self.projector = projector |
| 494 | self.feature_layer = feature_layer |
| 495 | |
| 496 | def forward(self, image): |
| 497 | all_hidden_states = self.tower( |
| 498 | image, output_hidden_states=True).hidden_states |
| 499 | features = all_hidden_states[self.feature_layer][:, 1:] |
| 500 | return self.projector(features) |
| 501 | |
| 502 | hf_config = AutoConfig.from_pretrained(args.model_path) |
| 503 | hf_config.vision_config._attn_implementation = "eager" |
| 504 | # Need to setup at hf_config._attn_implementation after transformers >= 4.46 |
| 505 | hf_config._attn_implementation = "eager" |
| 506 | model = LlavaForConditionalGeneration.from_pretrained( |
| 507 | args.model_path, dtype=torch.float16, config=hf_config) |
| 508 | wrapper = LlavaVisionWrapper( |
| 509 | model.vision_tower.to(args.device), |
| 510 | model.multi_modal_projector.to(args.device), |
| 511 | model.config.vision_feature_layer) |
| 512 | elif args.model_type == "llava_next": |
| 513 | from transformers import LlavaNextForConditionalGeneration |
| 514 | raw_image = Image.new('RGB', [512, 512]) |
| 515 | image = processor(text="dummy", images=raw_image, |
| 516 | return_tensors="pt")['pixel_values'].to( |
| 517 | args.device, torch.float16)[0] |
| 518 | |
| 519 | class LlavaNextVisionWrapper(torch.nn.Module): |
| 520 | |
| 521 | def __init__(self, vision_tower, projector): |
| 522 | super().__init__() |
| 523 | self.vision_tower = vision_tower |
| 524 | self.projector = projector |
| 525 | |
| 526 | def forward(self, pixel_values): |
| 527 | image_features = self.vision_tower(pixel_values, |
| 528 | output_hidden_states=True) |
| 529 | selected_image_feature = image_features.hidden_states[-2][:, 1:] |
| 530 | image_features = self.projector(selected_image_feature) |
| 531 | return image_features # (bs, 576, c) |
| 532 | |
| 533 | hf_config = AutoConfig.from_pretrained(args.model_path) |
| 534 | hf_config.vision_config._attn_implementation = "eager" |
| 535 | model = LlavaNextForConditionalGeneration.from_pretrained( |
| 536 | args.model_path, dtype=torch.float16, config=hf_config) |
| 537 | wrapper = LlavaNextVisionWrapper( |
no test coverage detected