mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
fix eos token not added correctly
This commit is contained in:
@@ -89,7 +89,8 @@ def get_prompts_tokens_with_weights(
|
||||
|
||||
def get_prompts_tokens_with_weights_t5(
|
||||
t5_tokenizer: T5Tokenizer,
|
||||
prompt: str
|
||||
prompt: str,
|
||||
add_special_tokens: bool = True
|
||||
):
|
||||
"""
|
||||
Get prompt token ids and weights, this function works for both prompt and negative prompt
|
||||
@@ -104,7 +105,7 @@ def get_prompts_tokens_with_weights_t5(
|
||||
inputs = t5_tokenizer(
|
||||
word,
|
||||
truncation=False, # so that tokenize whatever length prompt
|
||||
add_special_tokens=True,
|
||||
add_special_tokens=add_special_tokens,
|
||||
return_length=False,
|
||||
)
|
||||
|
||||
@@ -1458,20 +1459,21 @@ def get_weighted_text_embeddings_chroma(
|
||||
dtype = pipe.text_encoder.dtype
|
||||
|
||||
prompt_tokens, prompt_weights, prompt_masks = get_prompts_tokens_with_weights_t5(
|
||||
pipe.tokenizer, prompt
|
||||
pipe.tokenizer, prompt, add_special_tokens=False
|
||||
)
|
||||
|
||||
neg_prompt_tokens, neg_prompt_weights, neg_prompt_masks = get_prompts_tokens_with_weights_t5(
|
||||
pipe.tokenizer, neg_prompt
|
||||
pipe.tokenizer, neg_prompt, add_special_tokens=False
|
||||
)
|
||||
|
||||
padded_tokens, padded_weights, padded_masks = pad_prompt_tokens_to_same_size_chroma(
|
||||
pipe,
|
||||
[prompt_tokens, neg_prompt_tokens],
|
||||
[prompt_weights, neg_prompt_weights],
|
||||
[prompt_masks, neg_prompt_masks]
|
||||
[prompt_masks, neg_prompt_masks],
|
||||
add_eos_token=True
|
||||
)
|
||||
|
||||
|
||||
prompt_tokens = padded_tokens[0]
|
||||
prompt_weights = padded_weights[0]
|
||||
prompt_masks = padded_masks[0]
|
||||
@@ -1496,9 +1498,21 @@ def get_weighted_text_embeddings_chroma(
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
|
||||
# debug, will be removed later
|
||||
prompt_with_mask, prompt_without_mask = debug_masked_tokens(pipe, prompt_tokens, prompt_masks.detach().tolist()[0])
|
||||
neg_prompt_with_mask, neg_prompt_without_mask = debug_masked_tokens(pipe, neg_prompt_tokens, neg_prompt_masks.detach().tolist()[0])
|
||||
|
||||
return prompt_embeds, prompt_masks, neg_prompt_embeds, neg_prompt_masks
|
||||
|
||||
|
||||
# debug, will be removed later
|
||||
def debug_masked_tokens(pipe, prompt_tokens, prompt_masks):
|
||||
prompt_with_mask = pipe.tokenizer.decode([token for token, mask in zip(prompt_tokens, prompt_masks) if mask == 1], skip_special_tokens=False)
|
||||
prompt_without_mask = pipe.tokenizer.decode(prompt_tokens, skip_special_tokens=False)
|
||||
|
||||
return prompt_with_mask, prompt_without_mask
|
||||
|
||||
|
||||
def get_weighted_prompt_embeds_with_attention_mask_chroma(
|
||||
pipe: ChromaPipeline,
|
||||
tokens,
|
||||
@@ -1517,7 +1531,7 @@ def get_weighted_prompt_embeds_with_attention_mask_chroma(
|
||||
return prompt_embeds, prompt_masks
|
||||
|
||||
|
||||
def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, input_masks):
|
||||
def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, input_masks, min_length=3, add_eos_token=True):
|
||||
"""
|
||||
Implementation of Chroma's padding for prompt embeddings.
|
||||
Pads the embeddings to the maximum length found in the batch, while ensuring
|
||||
@@ -1531,13 +1545,14 @@ def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, inp
|
||||
input_masks = input_masks.copy()
|
||||
|
||||
pad_token_id = pipe.tokenizer.pad_token_id
|
||||
eos_token_id = pipe.tokenizer.eos_token_id
|
||||
|
||||
for tokens, mask in zip(input_tokens, input_masks):
|
||||
for j, token in enumerate(tokens):
|
||||
if token == pad_token_id:
|
||||
mask[j] = 0
|
||||
|
||||
max_token_count = max([len(x) for x in input_tokens])
|
||||
max_token_count = max([len(x) for x in input_tokens] + [min_length])
|
||||
|
||||
padded_tokens = []
|
||||
padded_weights = []
|
||||
@@ -1546,9 +1561,15 @@ def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, inp
|
||||
for tokens, weights, mask in zip(input_tokens, input_weights, input_masks):
|
||||
current_length = len(tokens)
|
||||
|
||||
pad_length = 0
|
||||
|
||||
if current_length < max_token_count:
|
||||
pad_length = max_token_count - current_length
|
||||
|
||||
elif pad_token_id not in tokens:
|
||||
pad_length = 1
|
||||
|
||||
if pad_length > 0:
|
||||
token_pad = [pad_token_id] * pad_length
|
||||
weight_pad = [1.0] * pad_length
|
||||
mask_pad = [0] * pad_length
|
||||
@@ -1561,6 +1582,8 @@ def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, inp
|
||||
padded_weights.append(weights)
|
||||
padded_masks.append(mask)
|
||||
|
||||
max_token_count = max([len(x) for x in padded_tokens])
|
||||
|
||||
for i, (tokens, weights, mask) in enumerate(zip(padded_tokens, padded_weights, padded_masks)):
|
||||
if pad_token_id in tokens:
|
||||
if tokens[-1] == pad_token_id:
|
||||
@@ -1572,12 +1595,20 @@ def pad_prompt_tokens_to_same_size_chroma(pipe, input_tokens, input_weights, inp
|
||||
padded_masks[i] = mask + [1]
|
||||
max_token_count = max(max_token_count, len(padded_tokens[i]))
|
||||
|
||||
for i, tokens in enumerate(padded_tokens):
|
||||
if len(tokens) < max_token_count:
|
||||
pad_length = max_token_count - len(tokens)
|
||||
tokens += [pad_token_id] * pad_length
|
||||
if add_eos_token:
|
||||
max_token_count += 1 # eos token
|
||||
|
||||
for i in range(len(padded_tokens)):
|
||||
if len(padded_tokens[i]) < max_token_count:
|
||||
pad_length = max_token_count - len(padded_tokens[i])
|
||||
padded_weights[i] += [1.0] * pad_length
|
||||
padded_masks[i][-1] = 0
|
||||
padded_masks[i] += [0] * (pad_length - 1) + [1]
|
||||
|
||||
if add_eos_token:
|
||||
padded_tokens[i] += [pad_token_id] * (pad_length - 1) + [eos_token_id]
|
||||
padded_masks[i][-2] = 1
|
||||
else:
|
||||
padded_tokens[i] += [pad_token_id] * pad_length
|
||||
|
||||
return padded_tokens, padded_weights, padded_masks
|
||||
|
||||
Reference in New Issue
Block a user