From f2f0cbccc98a5f463b72ddf1537fed0ca2c93a70 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Mon, 19 Aug 2024 16:53:05 +0800 Subject: [PATCH] Bug fix/encode prompt (#96) * update readme * fix bug in encode_prompt --- .../pipeline/pipeline_easyanimate_multi_text_encoder.py | 4 ++-- .../pipeline_easyanimate_multi_text_encoder_inpaint.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder.py b/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder.py index da2b8f4..ceb7e5c 100644 --- a/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder.py +++ b/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder.py @@ -312,7 +312,7 @@ class EasyAnimatePipeline_Multi_Text_Encoder(DiffusionPipeline): ) text_input_ids = text_inputs.input_ids if text_input_ids.shape[-1] > actual_max_sequence_length: - reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length]) + reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) text_inputs = tokenizer( reprompt, padding="max_length", @@ -379,7 +379,7 @@ class EasyAnimatePipeline_Multi_Text_Encoder(DiffusionPipeline): ) uncond_input_ids = uncond_input.input_ids if uncond_input_ids.shape[-1] > actual_max_sequence_length: - reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length]) + reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) uncond_input = tokenizer( reuncond_tokens, padding="max_length", diff --git a/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder_inpaint.py b/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder_inpaint.py index ddf8e42..bfe79c7 100644 --- a/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder_inpaint.py +++ b/easyanimate/pipeline/pipeline_easyanimate_multi_text_encoder_inpaint.py @@ -345,7 +345,7 @@ class EasyAnimatePipeline_Multi_Text_Encoder_Inpaint(DiffusionPipeline): ) text_input_ids = text_inputs.input_ids if text_input_ids.shape[-1] > actual_max_sequence_length: - reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length]) + reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) text_inputs = tokenizer( reprompt, padding="max_length", @@ -412,7 +412,7 @@ class EasyAnimatePipeline_Multi_Text_Encoder_Inpaint(DiffusionPipeline): ) uncond_input_ids = uncond_input.input_ids if uncond_input_ids.shape[-1] > actual_max_sequence_length: - reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length]) + reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) uncond_input = tokenizer( reuncond_tokens, padding="max_length",