
    ^(j*                        d dl mZ d dlZd dlZd dlZd dlZd dlmZ d dlm	Z	 e	rd dl
mZ d dlZd dlZd dlZ G d d      Z G d d	      Z G d
 d      ZdddZ edddg      ZdddZddZy)    )annotationsN)
namedtuple)TYPE_CHECKING)ModelPatcherc                  H    e Zd ZdZd	dZd
dZddZddZedd       Z	d Z
y)MultiGPUThreadPoola  Persistent thread pool for multi-GPU work distribution.

    Maintains one worker thread per extra GPU device. Each thread calls
    set_torch_device() once at startup so that compiled kernel caches
    (inductor/triton) stay warm across diffusion steps.
    c                h   g | _         i | _        i | _        |D ]  }t        j                         }t        j                         }|| j                  |<   || j                  |<   t        j                  | j                  |||fd      }|j                          | j                   j                  |        y )NT)targetargsdaemon)
_workers_work_queues_result_queuesqueueQueue	threadingThread_worker_loopstartappend)selfdevicesdevicewqrqts         8/Users/danicosta/Desktop/Flux2/ComfyUI/comfy/multigpu.py__init__zMultiGPUThreadPool.__init__   s    02=??AFBB(*Df%*,D'  (9(9R@PY]^AGGIMM  #     c                "   	 t         j                  j                  |       	 |j                         }|y |\  }}}	  ||i |}	|j                  |	d f       6# t        $ rL}t	        j
                  d| d|        	 |j                         }|Y d }~y |j                  d |f       +d }~ww xY w# t         j                  j                  $ r}|j                  d |f       Y d }~d }~wt        $ r}|j                  d |f       Y d }~d }~ww xY w)Nz)MultiGPUThreadPool: failed to set device z: )	comfymodel_managementset_torch_device	ExceptionloggingerrorgetputInterruptProcessingException)
r   r   work_qresult_qeitemfnr   kwargsresults
             r   r   zMultiGPUThreadPool._worker_loop&   s   		""33F; ::<D|#Bf(T,V,fd^,   	MMEfXRPQsSTzz|<dAY'	 	  ))FF (dAY'' (dAY''(s@   A B0 	B-!.B(B((B-0DC%%D1D		Dc                F    | j                   |   j                  |||f       y N)r   r(   )r   r   r.   r   r/   s        r   submitzMultiGPUThreadPool.submit>   s"    &!%%r4&89r   c                <    | j                   |   j                         S r2   )r   r'   )r   r   s     r   
get_resultzMultiGPUThreadPool.get_resultA   s    ""6*..00r   c                H    t        | j                  j                               S r2   )listr   keysr   s    r   r   zMultiGPUThreadPool.devicesD   s    D%%**,--r   c                    | j                   j                         D ]  }|j                  d         | j                  D ]  }|j	                  d        y )Ng      @)timeout)r   valuesr(   r   join)r   r   r   s      r   shutdownzMultiGPUThreadPool.shutdownH   sB    ##**,BFF4L -AFF3F r   N)r   list[torch.device])r   torch.devicer*   queue.Queuer+   rA   )r   r@   )returnr?   )__name__
__module____qualname____doc__r   r   r3   r5   propertyr   r>    r   r   r   r      s4    $(0:1 . . r   r   c                       e Zd ZddZd Zd Zy)
GPUOptionsc                     || _         || _        y r2   device_indexrelative_speed)r   rM   rN   s      r   r   zGPUOptions.__init__P   s    (,r   c                B    t        | j                  | j                        S r2   )rJ   rM   rN   r9   s    r   clonezGPUOptions.cloneT   s    $++T-@-@AAr   c                    d| j                   iS )NrN   )rN   r9   s    r   create_dictzGPUOptions.create_dictW   s    d11
 	
r   N)rM   intrN   float)rC   rD   rE   r   rP   rR   rH   r   r   rJ   rJ   O   s    -B
r   rJ   c                  (    e Zd Zd ZddZd ZddZy)GPUOptionsGroupc                    i | _         y r2   )optionsr9   s    r   r   zGPUOptionsGroup.__init__]   s	    .0r   c                6    || j                   |j                  <   y r2   )rX   rM   )r   infos     r   addzGPUOptionsGroup.add`   s    *.T&&'r   c                z    t               }| j                  j                         D ]  }|j                  |        |S r2   )rV   rX   r<   r[   )r   copts      r   rP   zGPUOptionsGroup.clonec   s1    <<&&(CEE#J )r   c                   i }|j                   g}|j                  d      D ]  }|j                  |j                           g }|D ]a  }| j                  j	                  |j
                  t        |j
                  d            }|j                         ||<   |j                  |       c t        |D cg c]  }|j                   c}      }	|j                         D ]  }
|
dxx   |	z  cc<    ||j                  d<   y c c}w )Nmultigpug      ?rL   rN   multigpu_options)load_deviceget_additional_models_with_keyr   rX   r'   indexrJ   rR   minrN   r<   model_options)r   model	opts_dictr   extra_modeldevice_opts_listr   device_optsx	min_speedvalues              r   registerzGPUOptionsGroup.registeri   s    	','8'8&9 ??
KKNN;223 L .0F,,**6<<QWQ]Q]nq9rsK + 7 7 9If##K0 
 3CD3Ca))3CDE	%%'E"#y0# (2;./ Es   2C>N)rZ   rJ   )rg   r   )rC   rD   rE   r   r[   rP   ro   rH   r   r   rV   rV   \   s    1/<r   rV   c                   | j                         } t               }| j                  d      }t        |      dkD  r"|D ]  }|j	                  |j
                          t        |      }t        j                  j                  d      }|D cg c]  }|| j
                  k7  s| }	}|	d|dz
   }
|
j                         }|D ]  }||v s|j                  |        t        |      dkD  r>|D ]
  }d}|rt        j                  j                         }|D ]  }|j                  |j
                  |k7  r |j                  | j                  k7  r:t        |dd      sH|j                         }t!        j"                  d|j                  j$                  j&                   d	|         n || j)                  |
      }d|_        | j                  d      }|j-                  |       | j/                  d|        | j1                          |
t3               }|j5                  |        nt!        j"                  d       t        |
      }|j	                  | j
                         | j                  d      }|D cg c]  }|j
                  |v s| }}t        |      t        |      k7  r"| j/                  d|       | j1                          | S c c}w c c}w )zSPrepare ModelPatcher to contain deepclones of its BaseModel and related properties.r`   r   F)exclude_currentN   is_multigpu_base_clonez%Reusing loaded multigpu deepclone of z for )new_load_deviceTzVNo extra torch devices need initialization, skipping initializing MultiGPU Work Units.)rP   setrc   lenr[   rb   r7   r!   r"   get_all_torch_devicescopyremoveloaded_modelsrg   clone_base_uuidgetattrr%   rZ   	__class__rC   deepclone_multigpurs   r   set_additional_modelsmatch_multigpu_clonesrV   ro   )rg   max_gpusgpu_optionsreuse_loadedskip_devicesmultigpu_modelsmmall_devicesdfull_extra_deviceslimit_extra_devicesextra_devicesskipr   device_patcherrz   lmallowed_devicesmnew_multigpu_modelss                       r   create_multigpu_deepclonesr   }   s   KKME5L:::FO
?a!BR^^, "%L
 ((>>u>UK%0K[A9J9J4J![K,[hqj9',,.M=   &  =A#F!N 5:4J4J4X4X4Z'Bxx' ~~/ ))U-B-BB "2'?G %'XXZNLL#HI]I]IgIgIpIpHqqvw}v~!  A ( %!&!9!9&!9!Q 59N1#BB:NO"">2''
OD; $< 	##%)+KU#mn -.O))*:::FO&5Zo/9Y1oZ
3#77##J0CD##%Lo Lf [s   K K8KKLoadBalancework_per_device	idle_timec                   | d   }t        | d   j                               }g }g }d}|j                         D ]
  }	||	d   z  } |D ]4  }
||
   d   }||z  |z  }|j                  |       |j                  |       6 t	        |      }i }t        ||      D ]
  \  }
}|||
<    |st        |d      S t        ||      D cg c]
  \  }}||z   }}}t        t        |      t        |      z
        }|r|||z  z  }t        ||      S c c}}w )zfOptimize work assigned to different devices, accounting for their relative speeds and splittable work.ra   multigpu_clonesg        rN   N)
r7   r8   r<   r   round_preservedzipr   absre   max)rf   
total_workreturn_idle_timework_normalizedrh   r   speed_per_devicer   total_speedoptsr   rN   relative_workdict_work_per_devicewrcompletion_timer   s                     r   load_balance_devicesr      sK   01I=!2388:;GOK  "t,-- # "6*+;<#N2kA/}-	  &o6O!$Wo!>'4V$ "?/66 '*/;K&LM&Lsqqs&LOMC(3+??@Ioj01	+Y77 Ns   D	c                L   | D cg c]  }t        |       }}t        |      }t        t        |             |z
  }t        |       D cg c]  \  }}||||   z
  f }}}|j	                  d d       t        |      D ]  }||   d   }||xx   dz  cc<    |S c c}w c c}}w )zBRound all values in a list, preserving the combined sum of values.c                    | d   S )Nrr   rH   )rl   s    r   <lambda>z!round_preserved.<locals>.<lambda>   s    !A$r   T)keyreverser   rr   )rS   sumround	enumeratesortrange)r<   rl   flooredtotal_floored	remainderi
fractionalrd   s           r   r   r      s      &&v!s1vvG&LMc&k"]2I09&0AB0A11a
l#0AJBOOO591a !  N '
 Cs   B	B )NF)rg   r   r   rS   r   rV   )FN)rf   z	dict[str]r   rS   r   rS   )r<   zlist[float])
__future__r   r   r   torchr%   collectionsr   typingr   comfy.model_patcherr   comfy.utilsr!   comfy.patcher_extensioncomfy.model_managementr   rJ   rV   r   r   r   r   rH   r   r   <module>r      so    "     "  0   <  < ~
 
< <BFR ):K(HI"8Hr   