From 93e2cf452399e21607ccee59352f9ff6d7a576de Mon Sep 17 00:00:00 2001 From: Bo Li Date: Wed, 12 Apr 2023 11:58:28 +0000 Subject: [PATCH 1/2] fix on med.py: ensure the prompt is also being repeated as the visual_embeds. --- lavis/models/med.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/lavis/models/med.py b/lavis/models/med.py index e963ffb3b..7e29d54f2 100644 --- a/lavis/models/med.py +++ b/lavis/models/med.py @@ -1331,6 +1331,9 @@ def generate_from_encoder( if not use_nucleus_sampling: num_beams = num_beams visual_embeds = visual_embeds.repeat_interleave(num_beams, dim=0) + tokenized_prompt.input_ids = tokenized_prompt.input_ids.repeat_interleave(num_beams, dim=0) + # Make sure that the prompt is repeated same number of times as the visual_embeds + assert visual_embeds.size(0) == tokenized_prompt.input_ids.size(0) image_atts = torch.ones(visual_embeds.size()[:-1], dtype=torch.long).to( self.device From 537a1defcac9e08a702ebd3a246edc9fc3a068c6 Mon Sep 17 00:00:00 2001 From: Bo Li Date: Wed, 12 Apr 2023 12:04:09 +0000 Subject: [PATCH 2/2] fix on med.py: add determine statement before repeat tokenized_prompt. --- lavis/models/med.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/lavis/models/med.py b/lavis/models/med.py index 7e29d54f2..9d969d614 100644 --- a/lavis/models/med.py +++ b/lavis/models/med.py @@ -1331,7 +1331,9 @@ def generate_from_encoder( if not use_nucleus_sampling: num_beams = num_beams visual_embeds = visual_embeds.repeat_interleave(num_beams, dim=0) - tokenized_prompt.input_ids = tokenized_prompt.input_ids.repeat_interleave(num_beams, dim=0) + + if visual_embeds.size(0) != tokenized_prompt.input_ids.size(0): + tokenized_prompt.input_ids = tokenized_prompt.input_ids.repeat_interleave(num_beams, dim=0) # Make sure that the prompt is repeated same number of times as the visual_embeds assert visual_embeds.size(0) == tokenized_prompt.input_ids.size(0)