| 365 | image = inputs['pixel_values'] |
| 366 | |
| 367 | class Blip2VisionWrapper(torch.nn.Module): |
| 368 | |
| 369 | def __init__(self, vision_model, qformer, projector, query_tokens): |
| 370 | super().__init__() |
| 371 | self.vision_model = vision_model |
| 372 | self.qformer = qformer |
| 373 | self.projector = projector |
| 374 | self.query_tokens = query_tokens |
| 375 | |
| 376 | def forward(self, image): |
| 377 | features = self.vision_model(image)[0] |
| 378 | qformer_output = self.qformer(query_embeds=self.query_tokens, |
| 379 | encoder_hidden_states=features, |
| 380 | return_dict=True) |
| 381 | return self.projector(qformer_output.last_hidden_state) |
| 382 | |
| 383 | model = Blip2ForConditionalGeneration.from_pretrained(args.model_path, |
| 384 | dtype=torch.float16) |
no outgoing calls
no test coverage detected