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

Class LlavaVisionWrapper

tensorrt_llm/tools/multimodal_builder.py:488–500  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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"

Callers 1

build_llava_engineFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected