MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / build_llava_engine

Function build_llava_engine

tensorrt_llm/tools/multimodal_builder.py:480–617  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

478
479
480def 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(

Callers 1

buildMethod · 0.85

Calls 11

LlavaVisionWrapperClass · 0.85
process_imagesFunction · 0.85
export_onnxFunction · 0.85
build_trt_engineFunction · 0.85
from_pretrainedMethod · 0.45
toMethod · 0.45
squeezeMethod · 0.45
get_modelMethod · 0.45

Tested by

no test coverage detected