(
self,
image_caption_processor: AutoProcessor,
image_caption_model: AutoModelForCausalLM,
llm_model: AutoModelForCausalLM,
llm_tokenizer: AutoTokenizer,
)
| 19 | |
| 20 | class PromptEnhancer(torch.nn.Module): |
| 21 | def __init__( |
| 22 | self, |
| 23 | image_caption_processor: AutoProcessor, |
| 24 | image_caption_model: AutoModelForCausalLM, |
| 25 | llm_model: AutoModelForCausalLM, |
| 26 | llm_tokenizer: AutoTokenizer, |
| 27 | ): |
| 28 | super().__init__() |
| 29 | self.image_caption_processor = image_caption_processor |
| 30 | self.image_caption_model = image_caption_model |
| 31 | self.llm_model = llm_model |
| 32 | self.llm_tokenizer = llm_tokenizer |
| 33 | self.device = image_caption_model.device |
| 34 | # model parameters and buffer sizes plus some extra 1GB. |
| 35 | self.model_size = ( |
| 36 | self.get_model_size(self.image_caption_model) |
| 37 | + self.get_model_size(self.llm_model) |
| 38 | + 1073741824 |
| 39 | ) |
| 40 | |
| 41 | def forward(self, prompt, image_conditioning, max_resulting_tokens): |
| 42 | enhanced_prompt = generate_cinematic_prompt( |
nothing calls this directly
no test coverage detected