a
    AdzA                     @  sJ  d dl mZ d dlZd dlmZ d dlmZ d dlZedZ	d-ddZ
ed	d
dgZG dd deZd.ddddZedZedZddddZG dd dZG dd dZd/ddddZG dd deZdd d!d"Zd#d$ Zdd d%d&Zed'ejZed(ejZd)d* Ze d+kr>d dl!Z!e!j"e!j#d, nd dl$Z$dS )0    )annotationsN)
namedtuple)Lista  
!start: (prompt | /[][():]/+)*
prompt: (emphasized | scheduled | alternate | plain | WHITESPACE)*
!emphasized: "(" prompt ")"
        | "(" prompt ":" prompt ")"
        | "[" prompt "]"
scheduled: "[" [prompt ":"] prompt ":" [WHITESPACE] NUMBER [WHITESPACE] "]"
alternate: "[" prompt ("|" [prompt])+ "]"
WHITESPACE: /\s+/
plain: /([^\\\[\]():|]|\\.)+/
%import common.SIGNED_NUMBER -> NUMBER
Fc                   sv   |du srdd|n|d|fdddd   fdd	fd
dt | D fdd| D S )a  
    >>> g = lambda p: get_learned_conditioning_prompt_schedules([p], 10)[0]
    >>> g("test")
    [[10, 'test']]
    >>> g("a [b:3]")
    [[3, 'a '], [10, 'a b']]
    >>> g("a [b: 3]")
    [[3, 'a '], [10, 'a b']]
    >>> g("a [[[b]]:2]")
    [[2, 'a '], [10, 'a [[b]]']]
    >>> g("[(a:2):3]")
    [[3, ''], [10, '(a:2)']]
    >>> g("a [b : c : 1] d")
    [[1, 'a b  d'], [10, 'a  c  d']]
    >>> g("a[b:[c:d:2]:1]e")
    [[1, 'abe'], [2, 'ace'], [10, 'ade']]
    >>> g("a [unbalanced")
    [[10, 'a [unbalanced']]
    >>> g("a [b:.5] c")
    [[5, 'a  c'], [10, 'a b c']]
    >>> g("a [{b|d{:.5] c")  # not handling this right now
    [[5, 'a  c'], [10, 'a {b|d{ c']]
    >>> g("((a][:b:c [d:3]")
    [[3, '((a][:b:c '], [10, '((a][:b:c d']]
    >>> g("[a|(b:1.1)]")
    [[1, 'a'], [2, '(b:1.1)'], [3, 'a'], [4, '(b:1.1)'], [5, 'a'], [6, '(b:1.1)'], [7, 'a'], [8, '(b:1.1)'], [9, 'a'], [10, '(b:1.1)']]
    >>> g("[fe|]male")
    [[1, 'female'], [2, 'male'], [3, 'female'], [4, 'male'], [5, 'female'], [6, 'male'], [7, 'female'], [8, 'male'], [9, 'female'], [10, 'male']]
    >>> g("[fe|||]male")
    [[1, 'female'], [2, 'male'], [3, 'male'], [4, 'male'], [5, 'female'], [6, 'male'], [7, 'male'], [8, 'male'], [9, 'female'], [10, 'male']]
    >>> g = lambda p: get_learned_conditioning_prompt_schedules([p], 10, 10)[0]
    >>> g("a [b:.5] c")
    [[10, 'a b c']]
    >>> g("a [b:1.5] c")
    [[5, 'a  c'], [10, 'a b c']]
    Nr         ?c                   s<   g G  fdddt j}| | tt S )Nc                      s.   e Zd Z fddZfddZdS )zVget_learned_conditioning_prompt_schedules.<locals>.collect_steps.<locals>.CollectStepsc                   s   |j d }t|}r,|dk r&| n|}nd|v rB|   }n| }tt||j d< |j d dkr||j d  d S )N   .)childrenfloatminintappend)selftreesv
flt_offset
int_offsetresstepsuse_old_scheduling =/var/www/html/stable-diffusion-webui/modules/prompt_parser.py	scheduledP   s    
z`get_learned_conditioning_prompt_schedules.<locals>.collect_steps.<locals>.CollectSteps.scheduledc                   s     tdd  d S Nr   )extendrange)r   r   r   r   r   r   	alternate^   s    z`get_learned_conditioning_prompt_schedules.<locals>.collect_steps.<locals>.CollectSteps.alternateN)__name__
__module____qualname__r   r   r   r   r   r   CollectStepsO   s   r#   )larkZVisitorvisitsortedset)r   r   r#   )r   r   r   r   r   collect_stepsL   s    z@get_learned_conditioning_prompt_schedules.<locals>.collect_stepsc                   s"   G  fdddt j}| |S )Nc                      s<   e Zd Z fddZ fddZdd Zdd Zd	d
 ZdS )zJget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStepc                 3  s(   |\}}}}} |kr|p dn|V  d S )Nr   r   )r   argsbeforeafter_whenstepr   r   r   f   s    zTget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.scheduledc                 3  s(   dd |D }| d t |  V  d S )Nc                 S  s   g | ]}|sd n|qS ) r   ).0argr   r   r   
<listcomp>j       zhget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.alternate.<locals>.<listcomp>r   )lenr   r)   r.   r   r   r   i   s    zTget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.alternatec                   s    fdd d  |S )Nc                 3  s.   t | tr| V  n| D ]} |E d H  qd S N)
isinstancestr)xgenflattenr   r   r=   m   s    
zaget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.start.<locals>.flattenr0   )joinr6   r   r<   r   startl   s    zPget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.startc                 s  s   |d j V  d S )Nr   )valuer6   r   r   r   plaint   s    zPget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.plainc                 s  s   |D ]
}|V  qd S r7   r   )r   datar	   metachildr   r   r   __default__v   s    zVget_learned_conditioning_prompt_schedules.<locals>.at_step.<locals>.AtStep.__default__N)r    r!   r"   r   r   r?   rA   rE   r   r.   r   r   AtStepe   s
   rF   )r$   Transformer	transform)r/   r   rF   r   r.   r   at_stepd   s    z:get_learned_conditioning_prompt_schedules.<locals>.at_stepc                   sJ   zt |  W n  tjjy.   | gg Y S 0  fdd D S )Nc                   s   g | ]}| |gqS r   r   )r1   t)rI   r   r   r   r3      r4   zSget_learned_conditioning_prompt_schedules.<locals>.get_schedule.<locals>.<listcomp>)schedule_parserparser$   
exceptionsZ	LarkError)prompt)rI   r(   r   )r   r   get_schedule{   s
    z?get_learned_conditioning_prompt_schedules.<locals>.get_schedulec                   s   i | ]}| |qS r   r   r1   rN   )rO   r   r   
<dictcomp>   r4   z=get_learned_conditioning_prompt_schedules.<locals>.<dictcomp>c                   s   g | ]} | qS r   r   rP   )
promptdictr   r   r3      r4   z=get_learned_conditioning_prompt_schedules.<locals>.<listcomp>)r'   )promptsZ
base_stepshires_stepsr   r   )rI   r(   r   rO   r   rR   r   r   r   )get_learned_conditioning_prompt_schedules   s    &
rU   ScheduledPromptConditioningend_at_stepcondc                      s"   e Zd ZdZd fdd	Z  ZS )SdConditioningz
    A list with prompts for stable diffusion's conditioner model.
    Can also specify width and height of created image - SDXL needs it.
    FNc                   sZ   t    | | |d u r |}|p.t|dd| _|p@t|dd | _|pRt|dd | _d S )Nis_negative_promptFwidthheight)super__init__r   getattrrZ   r[   r\   )r   rS   rZ   r[   r\   	copy_from	__class__r   r   r^      s    

zSdConditioning.__init__)FNNN)r    r!   r"   __doc__r^   __classcell__r   r   ra   r   rY      s   rY   zSdConditioning | list[str])rS   c                   s   g }t ||||}i }t||D ]\}}	||d}
|
durH||
 q tdd |	D |d}| |}g }t|	D ]F\ \}}t|tr fdd|	 D }n|  }|t
|| qt|||< || q |S )a  converts a list of prompts into a list of prompt schedules - each schedule is a list of ScheduledPromptConditioning, specifying the comdition (cond),
    and the sampling step at which this condition is to be replaced by the next one.

    Input:
    (model, ['a red crown', 'a [blue:green:5] jeweled crown'], 20)

    Output:
    [
        [
            ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886,  0.0229, -0.0523,  ..., -0.4901, -0.3066,  0.0674], ..., [ 0.3317, -0.5102, -0.4066,  ...,  0.4119, -0.7647, -1.0160]], device='cuda:0'))
        ],
        [
            ScheduledPromptConditioning(end_at_step=5, cond=tensor([[-0.3886,  0.0229, -0.0522,  ..., -0.4901, -0.3067,  0.0673], ..., [-0.0192,  0.3867, -0.4644,  ...,  0.1135, -0.3696, -0.4625]], device='cuda:0')),
            ScheduledPromptConditioning(end_at_step=20, cond=tensor([[-0.3886,  0.0229, -0.0522,  ..., -0.4901, -0.3067,  0.0673], ..., [-0.7352, -0.4356, -0.7888,  ...,  0.6994, -0.4312, -1.2593]], device='cuda:0'))
        ]
    ]
    Nc                 S  s   g | ]}|d  qS )r   r   r1   r:   r   r   r   r3      r4   z,get_learned_conditioning.<locals>.<listcomp>)r`   c                   s   i | ]\}}||  qS r   r   )r1   kr   ir   r   rQ      r4   z,get_learned_conditioning.<locals>.<dictcomp>)rU   zipgetr   rY   get_learned_conditioning	enumerater8   dictitemsrV   )modelrS   r   rT   r   r   Zprompt_schedulescacherN   Zprompt_schedulecachedtextscondscond_schedulerW   r,   rX   r   rg   r   rk      s&    


rk   z\bAND\bz7^((?:\s|.)*?)(?:\s*:\s*([-+]?(?:\d+\.?|\d*\.\d+)))?\s*$c                 C  s   g }i }t | }|  | D ]}t|}g }|D ]z}t|}|d urP| n|df\}	}
|
d urlt|
nd}
||	d }|d u rt	|}|
|	 |||	< |
||
f q2|
| q|||fS )Nr   )rY   clearre_ANDsplit	re_weightsearchgroupsr
   rj   r5   r   )rS   res_indexesprompt_indexesprompt_flat_listrN   Z
subpromptsindexesZ	subpromptmatchtextweightindexr   r   r   get_multicond_prompt_list   s&    


r   c                   @  s   e Zd ZdddZdS )%ComposableScheduledPromptConditioningr   c                 C  s   || _ || _d S r7   )	schedulesr   )r   r   r   r   r   r   r^      s    z.ComposableScheduledPromptConditioning.__init__N)r   r    r!   r"   r^   r   r   r   r   r      s   r   c                   @  s   e Zd Zdd ZdS )MulticondLearnedConditioningc                 C  s   || _ || _d S r7   shapebatch)r   r   r   r   r   r   r^      s    z%MulticondLearnedConditioning.__init__Nr   r   r   r   r   r      s   r   )returnc           
        sV   t |\}}}t| |||| g }|D ]}	| fdd|	D  q&tt|f|dS )aN  same as get_learned_conditioning, but returns a list of ScheduledPromptConditioning along with the weight objects for each prompt.
    For each prompt, the list is obtained by splitting the prompt using the AND separator.

    https://energy-based-model.github.io/Compositional-Visual-Generation-with-Composable-Diffusion-Models/
    c                   s   g | ]\}}t  | |qS r   )r   )r1   rh   r   Zlearned_conditioningr   r   r3   
  r4   z6get_multicond_learned_conditioning.<locals>.<listcomp>r   )r   rk   r   r   r5   )
ro   rS   r   rT   r   r{   r}   r|   r   r~   r   r   r   "get_multicond_learned_conditioning   s    r   c                      s(   e Zd Z fddZedd Z  ZS )DictWithShapec                   s   t    | | d S r7   )r]   r^   update)r   r:   r   ra   r   r   r^     s    
zDictWithShape.__init__c                 C  s
   | d j S )N	crossattnr   )r   r   r   r   r     s    zDictWithShape.shape)r    r!   r"   r^   propertyr   rd   r   r   ra   r   r     s   r   z'List[List[ScheduledPromptConditioning]]cc                   s    d d j }t|t}|rR|} fdd| D }t|t f|d j }n tjt f|j |j	|j
d}t D ]h\}}d}t|D ]\}	}
||
jkr|	} qq|r|| j  D ]\}}||| |< qqz|| j ||< qz|S )Nr   c                   s2   i | ]*\}}|t jt f|j |j|jd qS )devicedtype)torchzerosr5   r   r   r   )r1   rf   paramr   r   r   rQ     r4   z*reconstruct_cond_batch.<locals>.<dictcomp>r   r   )rX   r8   rm   rn   r   r5   r   r   r   r   r   rl   rW   )r   current_stepr   is_dictZ	dict_condr   rh   rt   target_indexcurrententryrf   r   r   r   reconstruct_cond_batch  s$    
 
r   c                 C  s   t dd | D }tt| D ]X}| | jd |kr| | dd  }||| | jd  dg}t| | |g| |< qt| S )Nc                 S  s   g | ]}|j d  qS )r   r   re   r   r   r   r3   7  r4   zstack_conds.<locals>.<listcomp>r   r   )maxr   r5   r   repeatr   vstackstack)tensorstoken_countrh   last_vectorlast_vector_repeatedr   r   r   stack_conds4  s    r   c                   s   | j d d jd j}g  g }| j D ]l}g }|D ]T}d}t|jD ]\}}	||	jkrB|} q^qB|t |jf  |j| j q0|| q$t d t	rt
 d  }
 fdd|
D }t||d j}nt j|j|jd}||fS )Nr   c                   s$   i | ]  t  fd dD qS )c                   s   g | ]}|  qS r   r   re   rf   r   r   r3   Y  r4   z:reconstruct_multicond_batch.<locals>.<dictcomp>.<listcomp>)r   )r1   r   r   r   rQ   Y  r4   z/reconstruct_multicond_batch.<locals>.<dictcomp>r   r   )r   r   rX   rl   rW   r   r5   r   r8   rm   listkeysr   r   r   tor   r   )r   r   r   Z
conds_listZcomposable_promptsZconds_for_batchZcomposable_promptr   r   r   r   stackedr   r   r   reconstruct_multicond_batchB  s(    

r   zZ
\\\(|
\\\)|
\\\[|
\\]|
\\\\|
\\|
\(|
\[|
:\s*([+-]?[.\d]+)\s*\)|
\)|
]|
[^\\()\[\]:]+|
:
z\s*\bBREAK\b\s*c                   s  g  g }g }d}d} fdd}t | D ]}|d} |d}| drf | dd d	g q*| d
kr~|t  q*| dkr|t  q*|dur|r|| t| q*| dkr|r|| | q*| dkr|r|| | q*t	t
| }t|D ]0\}	}
|	dkr" ddg  |
d	g qq*|D ]}||| q:|D ]}||| qPt dkrzdd	gg d}	|	d t k r |	 d  |	d  d krވ |	 d   |	d  d 7  <  |	d  n|	d7 }	q~ S )a  
    Parses a string with attention tokens and returns a list of pairs: text and its associated weight.
    Accepted tokens are:
      (abc) - increases attention to abc by a multiplier of 1.1
      (abc:3.12) - increases attention to abc by a multiplier of 3.12
      [abc] - decreases attention to abc by a multiplier of 1.1
      \( - literal character '('
      \[ - literal character '['
      \) - literal character ')'
      \] - literal character ']'
      \ - literal character ''
      anything else - just text

    >>> parse_prompt_attention('normal text')
    [['normal text', 1.0]]
    >>> parse_prompt_attention('an (important) word')
    [['an ', 1.0], ['important', 1.1], [' word', 1.0]]
    >>> parse_prompt_attention('(unbalanced')
    [['unbalanced', 1.1]]
    >>> parse_prompt_attention('\(literal\]')
    [['(literal]', 1.0]]
    >>> parse_prompt_attention('(unnecessary)(parens)')
    [['unnecessaryparens', 1.1]]
    >>> parse_prompt_attention('a (((house:1.3)) [on] a (hill:0.5), sun, (((sky))).')
    [['a ', 1.0],
     ['house', 1.5730000000000004],
     [' ', 1.1],
     ['on', 1.0],
     [' a ', 1.1],
     ['hill', 0.55],
     [', sun, ', 1.1],
     ['sky', 1.4641000000000006],
     ['.', 1.1]]
    g?g]tE?c                   s,   t | t D ]} | d  |9  < qd S r   )r   r5   )start_position
multiplierpr   r   r   multiply_range  s    z.parse_prompt_attention.<locals>.multiply_ranger   r   \Nr   ([)]ZBREAKr   r0   )re_attentionfinditergroup
startswithr   r5   popr
   rerw   re_breakrl   )r   Zround_bracketsZsquare_bracketsZround_bracket_multiplierZsquare_bracket_multiplierr   mr   partsrh   partposr   r   r   parse_prompt_attentions  sN    $




 r   __main__)Zoptionflags)NF)NF)NF)%
__future__r   r   collectionsr   typingr   r$   ZLarkrK   rU   rV   r   rY   rk   compilerv   rx   r   r   r   r   rm   r   r   r   r   Xr   Sr   r   r    doctesttestmodZNORMALIZE_WHITESPACEr   r   r   r   r   <module>   s8   


l0


Z
