mirror of
https://github.com/vladmandic/automatic
synced 2026-09-05 20:40:44 +02:00
Merge pull request #4460 from CalamitousFelicitousness/pixtralpromptfix
fix(prompt): handle multi-modal processors in token counter
This commit is contained in:
@@ -395,30 +395,39 @@ def get_prompt_schedule(prompt, steps):
|
||||
|
||||
def get_tokens(pipe, msg, prompt):
|
||||
global token_dict, token_type # pylint: disable=global-statement
|
||||
token_count = 0
|
||||
if shared.sd_loaded and hasattr(pipe, 'tokenizer') and pipe.tokenizer is not None:
|
||||
tokenizer = pipe.tokenizer
|
||||
# For multi-modal processors (e.g., PixtralProcessor), use the underlying text tokenizer
|
||||
if hasattr(tokenizer, 'tokenizer') and tokenizer.tokenizer is not None:
|
||||
tokenizer = tokenizer.tokenizer
|
||||
prompt = prompt.replace(' BOS ', ' !!!!!!!! ').replace(' EOS ', ' !!!!!!! ')
|
||||
debug(f'Prompt tokenizer: type={msg} prompt="{prompt}"')
|
||||
if token_dict is None or token_type != shared.sd_model_type:
|
||||
token_type = shared.sd_model_type
|
||||
fn = pipe.tokenizer.name_or_path
|
||||
fn = getattr(tokenizer, 'name_or_path', '')
|
||||
if fn.endswith('tokenizer'):
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'vocab.json')
|
||||
fn = os.path.join(fn, 'vocab.json')
|
||||
else:
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'tokenizer', 'vocab.json')
|
||||
fn = os.path.join(fn, 'tokenizer', 'vocab.json')
|
||||
token_dict = shared.readfile(fn, silent=True)
|
||||
for k, v in pipe.tokenizer.added_tokens_decoder.items():
|
||||
added_tokens = getattr(tokenizer, 'added_tokens_decoder', {})
|
||||
for k, v in added_tokens.items():
|
||||
token_dict[str(v)] = k
|
||||
shared.log.debug(f'Tokenizer: words={len(token_dict)} file="{fn}"')
|
||||
has_bos_token = pipe.tokenizer.bos_token_id is not None
|
||||
has_eos_token = pipe.tokenizer.eos_token_id is not None
|
||||
ids = pipe.tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
has_bos_token = getattr(tokenizer, 'bos_token_id', None) is not None
|
||||
has_eos_token = getattr(tokenizer, 'eos_token_id', None) is not None
|
||||
try:
|
||||
ids = tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
except Exception:
|
||||
ids = []
|
||||
if has_bos_token and has_eos_token:
|
||||
for i in range(len(ids)):
|
||||
if ids[i] == 21622:
|
||||
ids[i] = pipe.tokenizer.bos_token_id
|
||||
ids[i] = tokenizer.bos_token_id
|
||||
elif ids[i] == 15203:
|
||||
ids[i] = pipe.tokenizer.eos_token_id
|
||||
ids[i] = tokenizer.eos_token_id
|
||||
tokens = []
|
||||
for i in ids:
|
||||
try:
|
||||
|
||||
+13
-5
@@ -419,12 +419,20 @@ def update_token_counter(text):
|
||||
from modules import extra_networks
|
||||
prompt, _ = extra_networks.parse_prompt(text)
|
||||
if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None:
|
||||
has_bos_token = hasattr(shared.sd_model.tokenizer, 'bos_token_id') and shared.sd_model.tokenizer.bos_token_id is not None
|
||||
has_eos_token = hasattr(shared.sd_model.tokenizer, 'eos_token_id') and shared.sd_model.tokenizer.eos_token_id is not None
|
||||
ids = shared.sd_model.tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
tokenizer = shared.sd_model.tokenizer
|
||||
# For multi-modal processors (e.g., PixtralProcessor), use the underlying text tokenizer
|
||||
if hasattr(tokenizer, 'tokenizer') and tokenizer.tokenizer is not None:
|
||||
tokenizer = tokenizer.tokenizer
|
||||
has_bos_token = hasattr(tokenizer, 'bos_token_id') and tokenizer.bos_token_id is not None
|
||||
has_eos_token = hasattr(tokenizer, 'eos_token_id') and tokenizer.eos_token_id is not None
|
||||
try:
|
||||
ids = tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
except Exception:
|
||||
ids = []
|
||||
token_count = len(ids) - int(has_bos_token) - int(has_eos_token)
|
||||
max_length = shared.sd_model.tokenizer.model_max_length - int(has_bos_token) - int(has_eos_token)
|
||||
model_max_length = getattr(tokenizer, 'model_max_length', 0)
|
||||
max_length = model_max_length - int(has_bos_token) - int(has_eos_token)
|
||||
if max_length is None or max_length < 0 or max_length > 10000:
|
||||
max_length = 0
|
||||
return gr.update(value=f"<span class='gr-box gr-text-input'>{token_count}/{max_length}</span>", visible=token_count > 0)
|
||||
|
||||
Reference in New Issue
Block a user