§
    ‚ŠtjÛ‡  ã                  óh  — U d Z ddlmZ ddlZddl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 ddlmZ ddlmZ  ej        e¦  «        Z e¦   «         rddlZdd	lmZ dd
lmZmZmZmZmZmZ i Zded<   i Z ded<   i Z!ded<   ej"        d`d„¦   «         Z#ej"        dad„¦   «         Z$ej"        dbd„¦   «         Z%dbd„Z&dbd„Z'dcd„Z(ddd!„Z)ded#„Z*dcd$„Z+e,fZ-d%ed&<    e¦   «         re-ej.        ej/        ej0        ej1        fz  Z-dfd)„Z2dgdhd,„Z3did.„Z4djd/„Z5dkd4„Z6dld8„Z7dmd:„Z8d;Z9dnd?„Z:dodC„Z;i Z<dDedE<   dpdG„Z= e=dH¦  «        dqdJ„¦   «         Z> e=dK¦  «        dqdL„¦   «         Z? e=dMdN¦  «        dqdO„¦   «         Z@ e=dMdP¦  «        dqdQ„¦   «         ZAdqdR„ZBej"        drdT„¦   «         ZCdsdW„ZDdXZEdYZFdtd[„ZGdud]„ZHdsd^„ZIdsd_„ZJdS )vué  Shared export utilities used by all exporter backends.

Organised into five sections (search for the `# â”€â”€ Name â”€â”€` banners):

- **Patch and fix registries** â€” backend-keyed `_PATCHES` / `_FX_NODE_FIXES` /
  `_FX_PROGRAM_FIXES` populated via `@register_patch(backend, *paths)` /
  `@register_fx_node_fix` / `@register_fx_program_fix`, applied via
  `apply_patches` / `apply_fx_node_fixes` / `apply_fx_program_fixes`.
- **Recursive structure traversal** â€” internal helpers (`_map_leaf_tensors`,
  `_iter_leaf_tensors`) that drive every other tensor utility.
- **Public tensor utilities** â€” `get_leaf_tensors`, `duplicate_leaf_tensors`,
  `cast_leaf_tensors`, and `prepare_for_export` (sets attention/experts impl,
  patches non-exportable patterns, strips output flags).
- **Export input preparers** â€” `@register_export_input_preparer(marker)`
  registry that precomputes the per-encoder kwargs (`cu_seqlens`, `position_ids`,
  audio chunks, â€¦) the model would otherwise need data-dependent ops for.
- **Decomposition** â€” `decompose_prefill_decode` (split a generative forward
  into prefill + decode) and `decompose_multimodal` + `is_multimodal` (split a
  multimodal forward into one entry per submodule), backed by `_capture_forward`.
é    )ÚannotationsN)ÚMutableMapping)ÚAnyé   )Úlogging)Úis_torch_available)ÚPreTrainedModel)Ú'get_vision_bilinear_indices_and_weightsÚget_vision_cu_seqlensÚget_vision_merged_shapeÚget_vision_nearest_position_idsÚget_vision_position_idsÚget_vision_window_indexz*dict[str, list[tuple[Any, str, callable]]]Ú_PATCHESzdict[str, list[callable]]Ú_FX_NODE_FIXESÚ_FX_PROGRAM_FIXESÚobjr   Ú	attributeÚstrÚfactoryc              #  ó¶   K  — t          | |¦  «        }t          | | ||¦  «        ¦  «         	 dV — t          | ||¦  «         dS # t          | ||¦  «         w xY w)zNSwap `obj.<attribute>` with `factory(original)` for the duration of the block.N)ÚgetattrÚsetattr)r   r   r   Úoriginals       úZ/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/exporters/utils.pyÚpatch_attributer   Q   sq   è è € õ �s˜IÑ&Ô&€HÝˆC�˜G˜G HÑ-Ô-Ñ.Ô.Ð.ð*Øˆˆˆå��Y Ñ)Ô)Ð)Ð)Ð)ø���Y Ñ)Ô)Ð)Ð)øøøs   ®A ÁAÚpatchesúlist[tuple[Any, str, callable]]c           	   #  óÄ   K  — t          j        ¦   «         5 }| D ]*\  }}}|                     t          |||¦  «        ¦  «         Œ+dV — ddd¦  «         dS # 1 swxY w Y   dS )uñ   Install `(obj, attribute, factory)` patches for the duration of the block.

    Plural form of `patch_attribute` â€” each `factory(original)` returns the replacement
    callable. Originals are restored on exit, even if the body raises.
    N)Ú
contextlibÚ	ExitStackÚenter_contextr   )r   Ústackr   r   r   s        r   Úpatch_attributesr$   \   s¼   è è € õ 
Ô	Ñ	Ô	ð  5Ø'.ð 	Jð 	JÑ#ˆC�˜GØ×Ò¥°°YÀÑ HÔ HÑIÔIÐIÐIØˆˆˆðð ð ñ ô ð ð ð ð ð ð ð øøøð ð ð ð ð ð s   –2AÁAÁAÚbackendc              #  ó”   K  — t          t                               | g ¦  «        ¦  «        5  dV — ddd¦  «         dS # 1 swxY w Y   dS )z:Install `_PATCHES[backend]` for the duration of the block.N)r$   r   Úget)r%   s    r   Úapply_patchesr(   i   s‘   è è € õ 
�(Ÿ,š, w°Ñ3Ô3Ñ	4Ô	4ð ð Øˆˆˆðð ð ñ ô ð ð ð ð ð ð ð øøøð ð ð ð ð ð s   «=½AÁAc                ó   ‡ — ˆ fd„}|S )zKAppend the decorated `(gm, node) -> bool` fix to `_FX_NODE_FIXES[backend]`.c                ód   •— t                                ‰g ¦  «                             | ¦  «         | S ©N)r   Ú
setdefaultÚappend©Úfnr%   s    €r   Ú	decoratorz'register_fx_node_fix.<locals>.decorators   s,   ø€ Ý×!Ò! '¨2Ñ.Ô.×5Ò5°bÑ9Ô9Ð9Øˆ	ó    © ©r%   r0   s   ` r   Úregister_fx_node_fixr4   p   s$   ø€ ðð ð ð ð ð Ðr1   c                ó   ‡ — ˆ fd„}|S )u  Append the decorated `(exported_program) -> None` fix to `_FX_PROGRAM_FIXES[backend]`.

    Use this for fixes that need program-level context (range_constraints, graph_signature,
    state_dict) â€” the per-node `_FX_NODE_FIXES` shape only sees one node at a time.
    c                ód   •— t                                ‰g ¦  «                             | ¦  «         | S r+   )r   r,   r-   r.   s    €r   r0   z*register_fx_program_fix.<locals>.decorator�   s,   ø€ Ý×$Ò$ W¨bÑ1Ô1×8Ò8¸Ñ<Ô<Ð<Øˆ	r1   r2   r3   s   ` r   Úregister_fx_program_fixr7   z   s$   ø€ ðð ð ð ð ð Ðr1   ÚreturnÚNonec                óX   — t                                | g ¦  «        D ]} ||¦  «         ŒdS )zDApply `_FX_PROGRAM_FIXES[backend]` to `exported_program` (in place).N)r   r'   )r%   Úexported_programÚfixs      r   Úapply_fx_program_fixesr=   ˆ   s?   € å ×$Ò$ W¨bÑ1Ô1ð ð ˆØˆÐÑÔÐÐðð r1   Úpathsc                ó   ‡ ‡— ˆ ˆfd„}|S )um  Append the decorated `factory(original)` to `_PATCHES[backend]`, once per `path`.

    Each `path` is a dotted Python path like `"torch.where"`, `"torch.Tensor.unsqueeze"`,
    or `"transformers.models.nllb_moe.modeling_nllb_moe.NllbMoeTop2Router._cast_classifier"`.
    The rightmost segment is the attribute to swap; the rest is the object that owns it.
    Paths are resolved at decoration time â€” submodules are imported as needed, falling
    back to `getattr` for class attributes. A path that fails to resolve (e.g. the backend
    isn't installed) is silently skipped so the module still imports.

    Passing multiple paths registers the SAME factory against each â€” useful for swapping
    the same method or torch op across several call sites (e.g. ``torch.unsqueeze`` +
    ``torch.Tensor.unsqueeze``, or one vision-attention forward across N model classes).
    c                óÊ   •— ‰D ]^}|                      d¦  «        \  }}}t          |¦  «        }|€Œ-t                               ‰g ¦  «                             ||| f¦  «         Œ_| S )Nú.)Ú
rpartitionÚ_resolve_dotted_pathr   r,   r-   )r/   ÚpathÚobj_pathÚ_r   r   r%   r>   s         €€r   r0   z!register_patch.<locals>.decorator�   st   ø€ Øð 	Jð 	JˆDØ%)§_¢_°SÑ%9Ô%9Ñ"ˆH�a˜Ý& xÑ0Ô0ˆCØˆ{ØÝ×Ò ¨Ñ,Ô,×3Ò3°S¸)ÀRÐ4HÑIÔIÐIÐIØˆ	r1   r2   )r%   r>   r0   s   `` r   Úregister_patchrG   Ž   s*   øø€ ðð ð ð ð ð ð Ðr1   rD   c                óH  — ddl }|                      d¦  «        }	 |                     |d         ¦  «        }|dd…         D ]I}	 |                     |j        › d|› �¦  «        }Œ## t          t
          f$ r t          ||¦  «        }Y ŒFw xY w|S # t          t
          f$ r Y dS w xY w)uù   Resolve a dotted Python path to the actual object â€” importing submodules where
    possible, falling back to `getattr` for class attributes (e.g. `torch.Tensor`).
    Returns `None` if the path can't be resolved (e.g. the backend isn't installed).r   NrA   é   )Ú	importlibÚsplitÚimport_moduleÚ__name__ÚImportErrorÚAttributeErrorr   )rD   rJ   Úpartsr   Úparts        r   rC   rC   ©   sÛ   € ð ÐÐÐà�JŠJ�s‰OŒO€Eð	Ø×%Ò% e¨A¤hÑ/Ô/ˆØ˜!˜"˜"”Ið 	)ð 	)ˆDð)Ø×-Ò-°´Ð.FÐ.FÀÐ.FÐ.FÑGÔG��øÝ¥Ð0ð )ð )ð )Ý˜c 4Ñ(Ô(���ð)øøøàˆ
øÝ�Ð(ð ð ð Øˆtˆtðøøøs5   ›'B ÁA#Á"B Á#!BÂB ÂBÂB ÂB!Â B!c                óª  — t                                | g ¦  «        }|                     ¦   «         D ]¢}t          |t          j        j        ¦  «        sŒ"t          |j        j	        ¦  «        D ]!}|j
        dk    rŒ|D ]} |||¦  «        r nŒŒ"	 |j                             ¦   «          |                     ¦   «          ŒŒ# t          t          f$ r Y ŒŸw xY wdS )u²  Walk every call_function node and apply the first matching `_FX_NODE_FIXES[backend]`
    fix, then DCE.

    Each fix has signature `(gm, node) -> bool`. Returning `True` means the fix consumed
    the node â€” no further fixes run against it. Fixes are expected to be disjoint by
    `node.target`; if multiple could apply, list order decides.

    After the walk, `Graph.eliminate_dead_code` runs on every sub-GraphModule and
    `gm.recompile()` is called once. PyTorch DCE occasionally raises `SystemError` /
    `KeyError` from `erase_node._update_args_kwargs` on orphaned symbolic-size nodes â€”
    we swallow both; any survivors are handled by the downstream backend optimizer.
    Úcall_functionN)r   r'   ÚmodulesÚ
isinstanceÚtorchÚfxÚGraphModuleÚlistÚgraphÚnodesÚopÚeliminate_dead_codeÚ	recompileÚSystemErrorÚKeyError)r%   Úgraph_moduleÚfixesÚgmÚnoder<   s         r   Úapply_fx_node_fixesre   ¼   sü   € õ ×Ò˜w¨Ñ+Ô+€EØ×"Ò"Ñ$Ô$ð ð ˆÝ˜"�eœhÔ2Ñ3Ô3ð 	ØÝ˜œœÑ(Ô(ð 	ð 	ˆDØŒw˜/Ò)Ð)ØØð ð �Ø�3�r˜4‘=”=ð Ø�Eðøð	ØŒH×(Ò(Ñ*Ô*Ð*Ø�LŠL‰NŒNˆNˆNøÝ�XÐ&ð 	ð 	ð 	ØˆDð	øøøðð s   Â-B<Â<CÃCztuple[type, ...]Ú_LEAF_SKIP_TYPESr/   Úcallablec           	     ó<  ‡— t          | t          ¦  «        r| S t          | t          j        ¦  «        r ‰| ¦  «        S t          | t          t
          t          f¦  «        r$ t          | ¦  «        ˆfd„| D ¦   «         ¦  «        S t          | t          ¦  «        r-t	          | ¦  «        D ]}t          | |         ‰¦  «        | |<   Œ| S t          | d¦  «        rFt          | ¦  «                             ¦   «         D ]$\  }}t          | |t          |‰¦  «        ¦  «         Œ%| S )u•  Apply `fn` to every tensor in a nested structure, preserving container types.

    Mutates dicts and `__dict__`-bearing objects in place (preserving identity â€” callers
    rely on this so downstream pops/mutations propagate back to the original mapping);
    rebuilds lists/tuples/sets/frozensets (immutable or order-sensitive containers).
    Skips non-traversable leaf types (enum, SymInt, etc.).
    c              3  ó8   •K  — | ]}t          |‰¦  «        V — Œd S r+   ©Ú_map_leaf_tensors)Ú.0Úitemr/   s     €r   ú	<genexpr>z$_map_leaf_tensors.<locals>.<genexpr>ó   s.   øè è € ÐEÐE¸Õ*¨4°Ñ4Ô4ÐEÐEÐEÐEÐEÐEr1   Ú__dict__)rU   rf   rV   ÚTensorrY   ÚtupleÚsetÚtypeÚdictrk   ÚhasattrÚvarsÚitemsr   )r   r/   ÚkÚattrÚattr_vals    `   r   rk   rk   æ   s%  ø€ õ �#Õ'Ñ(Ô(ð Øˆ
Ý�#•u”|Ñ$Ô$ð Øˆr�#‰wŒwˆÝ�#��e¥SÐ)Ñ*Ô*ð FØ�t�C‰yŒyÐEÐEÐEÐEÀÐEÑEÔEÑEÔEÐEÝ�#•tÑÔð Ý�c‘”ð 	3ð 	3ˆAÝ& s¨1¤v¨rÑ2Ô2ˆC�‰FˆFØˆ
Ýˆs�JÑÔð @Ý" 3™iœiŸošoÑ/Ô/ð 	@ð 	@‰NˆD�(Ý�C˜Õ0°¸2Ñ>Ô>Ñ?Ô?Ð?Ð?Ø€Jr1   Ú Úprefixc              #  ó\  K  — t          | t          ¦  «        rdS t          | t          j        ¦  «        r
|pd| fV — dS t          | t          t
          t          f¦  «        rEt          | ¦  «        D ]3\  }}|r|› d|› �nt          |¦  «        }t          ||¦  «        E d{V —† Œ4dS t          | t          ¦  «        r=|                      ¦   «         D ]&\  }}|r|› d|› �n|}t          ||¦  «        E d{V —† Œ'dS t          | d¦  «        r%t          t          | ¦  «        |¦  «        E d{V —† dS dS )zEYield `(dotted_path, tensor)` for every tensor in a nested structure.NÚoutputrA   ro   )rU   rf   rV   rp   rY   rq   rr   Ú	enumerater   Ú_iter_leaf_tensorsrt   rw   ru   rv   )r   r|   Úindexrm   rD   ÚkeyÚvalues          r   r€   r€   þ   s™  è è € å�#Õ'Ñ(Ô(ð ØˆÝ�#•u”|Ñ$Ô$ð 9ØÐ ˜ #Ð%Ð%Ð%Ð%Ð%Ð%Ý	�C�$¥¥sÐ+Ñ	,Ô	,ð 	9Ý$ S™>œ>ð 	6ð 	6‰KˆE�4Ø*0Ð@�fÐ&Ð&˜uÐ&Ð&Ð&µc¸%±j´jˆDÝ)¨$°Ñ5Ô5Ð5Ð5Ð5Ð5Ð5Ð5Ð5Ð5ð	6ð 	6õ 
�C�Ñ	Ô	ð 9ØŸ)š)™+œ+ð 	7ð 	7‰JˆC�Ø(.Ð7�fÐ$Ð$˜sÐ$Ð$Ð$°CˆDÝ)¨%°Ñ6Ô6Ð6Ð6Ð6Ð6Ð6Ð6Ð6Ð6ð	7ð 	7õ 
��jÑ	!Ô	!ð 9Ý%¥d¨3¡i¤i°Ñ8Ô8Ð8Ð8Ð8Ð8Ð8Ð8Ð8Ð8Ð8ð9ð 9r1   údict[str, torch.Tensor]c                ó:   — t          t          | ¦  «        ¦  «        S )a  Recursively retrieve all leaf tensors from a potentially nested structure.

    Args:
        obj (`Any`):
            A tensor, dataclass, dict, list, tuple, or any nesting thereof.

    Returns:
        `dict[str, torch.Tensor]`: Flat mapping from dotted path strings to tensors.
    )rt   r€   )r   s    r   Úget_leaf_tensorsr†     s   € õ Õ" 3Ñ'Ô'Ñ(Ô(Ð(r1   c                óL   ‡— t          ¦   «         Šdˆfd„}t          | |¦  «        S )a”  Clone tensors that appear more than once in an output structure.

    When a model returns the same tensor under two output names (e.g. `last_hidden_state`
    and `hidden_states[0]`), the ONNX optimizer deduplicates the two output nodes and
    renames one, breaking the expected name mapping. Cloning duplicates gives each output
    leaf a distinct identity so the optimizer has nothing to merge.
    Útensorútorch.Tensorr8   c                ó–   •— t          | ¦  «        ‰v r|                      ¦   «         S ‰                     t          | ¦  «        ¦  «         | S r+   )ÚidÚcloneÚadd)rˆ   Úseens    €r   Ú_dedupz&duplicate_leaf_tensors.<locals>._dedup+  s?   ø€ Ýˆf‰:Œ:˜ÐÐØ—<’<‘>”>Ð!Ø�Š•�F‘”ÑÔÐØˆr1   ©rˆ   r‰   r8   r‰   )rr   rk   )r   r�   rŽ   s     @r   Úduplicate_leaf_tensorsr‘   !  s>   ø€ õ ‰5Œ5€Dðð ð ð ð ð õ ˜S &Ñ)Ô)Ð)r1   Údtypeútorch.dtypeÚdeviceútorch.devicec                ó4   ‡‡— dˆˆfd„}t          | |¦  «        S )zJRecursively cast all floating-point tensors to the given dtype and device.rˆ   r‰   r8   c                ó†   •— |                       ¦   «         r|                      ‰‰¬¦  «        n|                      ‰¬¦  «        S )N©r’   r”   )r”   )Úis_floating_pointÚto)rˆ   r”   r’   s    €€r   Ú_castz cast_leaf_tensors.<locals>._cast7  sB   ø€ Ø8>×8PÒ8PÑ8RÔ8RÐpˆv�yŠy˜u¨VˆyÑ4Ô4Ð4ÐX^×XaÒXaÐioÐXaÑXpÔXpÐpr1   r�   rj   )r   r’   r”   r›   s    `` r   Úcast_leaf_tensorsrœ   4  s@   øø€ ðqð qð qð qð qð qð qõ ˜S %Ñ(Ô(Ð(r1   Úmodelú!PreTrainedModel | torch.nn.Moduleútorch.device | Nonec                ó    — t          | d¦  «        r| j        S 	 t          |                      ¦   «         ¦  «        j        S # t          $ r Y dS w xY w)a#  `.device` for any `nn.Module`. `PreTrainedModel` exposes it directly via `ModuleUtilsMixin`;
    for plain submodules (e.g. a `Linear` or `MultiModalProjector` from a decomposed multimodal model)
    we fall back to the first parameter. Returns `None` if the module has no parameters at all.r”   N)ru   r”   ÚnextÚ
parametersÚStopIteration©r�   s    r   Úmodule_devicer¥   =  sb   € õ ˆu�hÑÔð ØŒ|ÐðÝ�E×$Ò$Ñ&Ô&Ñ'Ô'Ô.Ð.øÝð ð ð Øˆtˆtðøøøó   ™%? ¿
AÁAútorch.dtype | Nonec                ó    — t          | d¦  «        r| j        S 	 t          |                      ¦   «         ¦  «        j        S # t          $ r Y dS w xY w)zE`.dtype` for any `nn.Module`. Same fallback story as `module_device`.r’   N)ru   r’   r¡   r¢   r£   r¤   s    r   Úmodule_dtyper©   I  s`   € åˆu�gÑÔð ØŒ{ÐðÝ�E×$Ò$Ñ&Ô&Ñ'Ô'Ô-Ð-øÝð ð ð Øˆtˆtðøøør¦   )Ú	use_cacheÚoutput_attentionsÚoutput_hidden_statesÚreturn_dictÚreturn_lossÚinputsúMutableMapping[str, Any]úRtuple[PreTrainedModel | torch.nn.Module, MutableMapping[str, Any], dict[str, Any]]c                ó,  ‡— dD ]0}‰                      |d¦  «        }|�t          d|› d|› d�¦  «        ‚Œ1t          | d¦  «        r%t          | j        dd¦  «        rt          d	¦  «        ‚‰                     dd¦  «        rt          d
¦  «        ‚ˆfd„t          D ¦   «         }t          j        ¦   «         5  t          | ‰¦  «         ddd¦  «         n# 1 swxY w Y   t          | ¦  «        }t          | ¦  «        }|€|�t          ‰||¬¦  «        Š| ‰|fS )u6  Configure model and inputs for export. Mutates both `model` and `inputs` in place,
    returning `(model, inputs, output_flags)` where `output_flags` holds the values popped
    from `inputs` for `use_cache`, `return_dict`, etc. (to be applied reversibly onto
    `model.config` by `patch_model_config` during the trace).

    - Strips label inputs (`labels`, `future_values`) â€” loss computation is unsupported.
    - Pops output flags (`use_cache`, `return_dict`, â€¦) from `inputs` so they don't appear
      as traced kwargs; the values are returned for the trace block to apply onto
      `model.config`.
    - Pre-computes data-dependent vision/audio kwargs registered via
      `@register_export_input_preparer` and writes them into `inputs`.
    - Casts input tensors to match the model's `dtype` / `device`.
    )ÚlabelsÚfuture_valuesNzFound 'zM' in inputs. Loss computation is not supported during export. Please remove 'z+' from your inputs before calling export().Úconfigr®   FzœFound 'model.config.return_loss=True'. Loss computation is not supported during export. Please set 'model.config.return_loss=False' before calling export().z•Found 'return_loss=True' in inputs. Loss computation is not supported during export. Please remove 'return_loss' from your inputs or set it to False.c                óD   •— i | ]}|‰v ¯|‰                      |¦  «        “ŒS r2   )Úpop)rl   Úflagr¯   s     €r   ú
<dictcomp>z&prepare_for_export.<locals>.<dictcomp>|  s-   ø€ ÐWÐWÐW¨tÈÐPVÈÈ�D˜&Ÿ*š* TÑ*Ô*ÈÈÈr1   r˜   )r·   Ú
ValueErrorru   r   rµ   r'   Ú_OUTPUT_FLAGSrV   Úno_gradÚprecompute_export_inputsr©   r¥   rœ   )r�   r¯   Ú	label_keyrƒ   Úoutput_flagsr’   r”   s    `     r   Úprepare_for_exportrÀ   W  sµ  ø€ ð" 1ð ð ˆ	Ø—
’
˜9 dÑ+Ô+ˆØÐÝðY˜)ð Yð YØ"+ðYð Yð Yñô ð ð õ
 ˆu�hÑÔð 
¥G¨E¬L¸-ÈÑ$OÔ$Oð 
ÝðSñ
ô 
ð 	
ð ‡z‚z�- Ñ'Ô'ð 
ÝðOñ
ô 
ð 	
ð XÐWÐWÐWµ}ÐWÑWÔW€Lõ
 
Œ‰Œð 0ð 0Ý  ¨Ñ/Ô/Ð/ð0ð 0ð 0ñ 0ô 0ð 0ð 0ð 0ð 0ð 0ð 0øøøð 0ð 0ð 0ð 0õ
 ˜ÑÔ€EÝ˜5Ñ!Ô!€FØÐ˜FÐ.Ý" 6°¸vÐFÑFÔFˆà�&˜,Ð&Ð&s   Â5CÃCÃCútorch.nn.ModuleÚnameú
Any | Nonec                ób   — |                       ¦   «         D ]}t          ||d¦  «        x}�|c S ŒdS )zTReturn the first non-None value of `name` found on `model` or any of its submodules.N)rT   r   )r�   rÂ   Úmodulerƒ   s       r   Ú_find_submodule_attrrÆ   •  sC   € à—-’-‘/”/ð ð ˆÝ˜V T¨4Ñ0Ô0Ð0ˆEÐ=ØˆLˆLˆLð >àˆ4r1   zdict[tuple[str, ...], callable]Ú_EXPORT_INPUT_PREPARERSÚmarkersc                 ó   ‡ — ˆ fd„}|S )u4  Register `fn(model, inputs) -> None`. Dispatched when every `marker` is a key in
    `inputs` with a non-`None` value â€” no model_type list to maintain. Use multiple
    markers to narrow the match when a single kwarg is too ambiguous (e.g.
    `("input_features", "feature_lens")` for omni audio encoders).c                ó   •— | t           ‰<   | S r+   )rÇ   )r/   rÈ   s    €r   r0   z1register_export_input_preparer.<locals>.decorator¦  s   ø€ Ø+-Õ Ñ(Øˆ	r1   r2   )rÈ   r0   s   ` r   Úregister_export_input_preparerrË      s$   ø€ ðð ð ð ð ð Ðr1   Úgrid_thwúdict[str, Any]c                ó²  — |d         }t          | d¦  «        }|€|                     dd¦  «        }t          |¦  «        |d<   t          | d¦  «        du}t          |||¬¦  «        |d	<   t          | d
¦  «        }t          | d¦  «        }|�|�t	          ||||¦  «        \  |d<   |d<   t          | d¦  «        }|�t          |||¦  «        \  |d<   |d<   dS dS )aã  Precompute helpers driven by `grid_thw`: `cu_seqlens`, `position_ids`, plus optional
    `window_index`/`cu_window_seqlens` (XNet-style window attn) and
    `bilinear_indices`/`bilinear_weights` (interpolation-based merging).

    Optional helpers are gated by the presence of their config attribute on the encoder
    (`window_size`+`patch_size` for window attention, `num_grid_per_side` for bilinear),
    so a model that doesn't use that feature won't get its kwarg injected.
    rÌ   Úspatial_merge_sizeNÚmerge_sizesrI   Ú
cu_seqlensÚaxis_dim)Úinclude_temporalÚposition_idsÚwindow_sizeÚ
patch_sizeÚwindow_indexÚcu_window_seqlensÚnum_grid_per_sideÚbilinear_indicesÚbilinear_weights)rÆ   r'   r   r   r   r
   )r�   r¯   rÌ   rÏ   rÓ   rÕ   rÖ   rÙ   s           r   Ú_prepare_grid_thw_vision_inputsrÜ   ­  s#  € ð �jÔ!€HÝ-¨eÐ5IÑJÔJÐØÐ!ð $ŸZšZ¨°qÑ9Ô9Ðå0°Ñ:Ô:€Fˆ<Ñõ ,¨E°:Ñ>Ô>ÀdÐJÐÝ4°XÐ?QÐdtÐuÑuÔu€Fˆ>Ñå& u¨mÑ<Ô<€KÝ% e¨\Ñ:Ô:€JØÐ :Ð#9Ý>UØÐ(¨+°zñ?
ô ?
Ñ;ˆˆ~Ñ Ð':Ñ ;õ -¨UÐ4GÑHÔHÐØÐ$ÝAhØÐ'Ð);ñB
ô B
Ñ>ˆÐ!Ñ" FÐ+=Ñ$>Ð$>Ð$>ð %Ð$r1   Útarget_sizesc                ó@  — |d         }t          | d¦  «        }|�t          ||¦  «        |d<   t          | d¦  «        }|�^t          j        j                             |dd¬¦  «        }t          |d|d	         d¬
¦  «        \  |d<   |d<   t          ||¦  «        |d<   dS dS )a
  NaViT-style packed encoders carry per-image `(h, w)` as `target_sizes` instead of `grid_thw`.
    Synthesise `grid_thw = [1, h, w]` and run the nearest-position-id / window-index /
    merged-shape helpers so the per-image Python loops move outside the traced graph.rÝ   Únum_patches_per_sideNrÔ   Úwindow_kernel_size)rI   r   rI   )rƒ   r   )rÏ   rÕ   rÖ   r×   rØ   Úmerged_shape)rÆ   r   rV   ÚnnÚ
functionalÚpadr   r   )r�   r¯   rÝ   rß   rà   rÌ   s         r   Ú_prepare_navit_vision_inputsrå   Ò  sÉ   € ð
 ˜.Ô)€LÝ/°Ð7MÑNÔNÐØÐ'Ý!@ÀÐOcÑ!dÔ!dˆˆ~Ñå-¨eÐ5IÑJÔJÐØÐ%Ý”8Ô&×*Ò*¨<¸ÀqÐ*ÑIÔIˆÝ>UØ¨Ð8JÈ1Ô8MÐZ[ð?
ñ ?
ô ?
Ñ;ˆˆ~Ñ Ð':Ñ ;õ "9¸ÐGYÑ!ZÔ!Zˆˆ~ÑÐÐð &Ð%r1   Úinput_featuresÚfeature_lensc                óþ  — |d         }|d         }t           j        t          | ¦  «        j                 }t	          |d¦  «        }t	          |d¦  «        }t	          |d¦  «        } |||| j        ¦  «        \  }}	||d<   |	|d<   t          | d¦  «        r1 ||	|| j        | j        ¦  «        |d	<    ||	| j        ¦  «        |d
<   dS  ||	¦  «        |d	<    ||	¦  «        |d
<    t	          |d¦  «        |¦  «        |d<   dS )uŽ  Replace `input_features`/`feature_lens` with precomputed `padded_feature`, `chunk_lengths`,
    `cu_seqlens`, `valid_indices` (+ `pool_indices` on Qwen2.5-Omni-style encoders) so the
    encoder's `.split(.tolist(), dim=0)` and related data-dependent ops happen outside the
    traced graph.

    The helpers (`chunk_and_pad_features`, `get_audio_cu_seqlens`, â€¦) all live in the model's
    own ``modeling_*.py`` module, so we resolve them via ``type(model).__module__`` rather than
    hard-coding one Omni variant. ``n_window_infer`` selects the Qwen3-Omni-style four-arg
    ``get_audio_cu_seqlens`` over the Qwen2.5-Omni-style single-arg form.
    rç   ræ   Úchunk_and_pad_featuresÚget_audio_cu_seqlensÚget_valid_indicesÚpadded_featureÚchunk_lengthsÚn_window_inferrÑ   Úvalid_indicesÚget_pool_indicesÚpool_indicesN)ÚsysrT   rs   Ú
__module__r   Ún_windowru   rî   )
r�   r¯   rç   ræ   rÅ   ré   rê   rë   rì   rí   s
             r   Ú_prepare_omni_audio_inputsrõ   å  s8  € ð ˜.Ô)€LØÐ,Ô-€NÝŒ[�˜e™œÔ/Ô0€Få$ VÐ-EÑFÔFÐÝ" 6Ð+AÑBÔBÐÝ Ð(;Ñ<Ô<Ðà$:Ð$:¸>È<ÐY^ÔYgÑ$hÔ$hÑ!€N�MØ-€FÐÑØ+€Fˆ?ÑÝˆuÐ&Ñ'Ô'ð SØ3Ð3°MÀ<ÐQVÔQeÐglÔguÑvÔvˆˆ|ÑØ"3Ð"3°MÀ5Ä>Ñ"RÔ"RˆˆÑÐÐà3Ð3°MÑBÔBˆˆ|ÑØ"3Ð"3°MÑ"BÔ"BˆˆÑØ!D¥¨Ð1CÑ!DÔ!DÀ\Ñ!RÔ!Rˆˆ~ÑÐÐr1   Úinput_features_maskc                óÎ  — ddl m} t          | d¦  «        }t          | d¦  «        }|�|€dS |d         }|j        \  }}||dz  z  }|                     d¦  «                             t          j        ¦  «        }	|                     ||d¦  «                             d¬¦  «         	                    d¦  «                             t          j        ¦  «        }
 ||
|	||¦  «        |d	<   dS )
u4  Precompute `cu_seqlens` for Qwen3-ASR â€” the encoder's call to ``get_audio_cu_seqlens``
    has a data-dependent Python loop that we evaluate here so the encoder pops the result
    from ``kwargs``. Mirrors the few lines that build ``feature_lens``/``chunk_lengths`` in
    ``Qwen3ASREncoder.forward``.
    r   )rê   rô   rî   Nrö   éÿÿÿÿ)ÚdimrÑ   )
Ú#models.qwen3_asr.modeling_qwen3_asrrê   rÆ   ÚshapeÚsumrš   rV   ÚlongÚviewÚreshape)r�   r¯   rê   rô   rî   rö   Ú
batch_sizeÚpadded_feature_lengthÚ
num_chunksrç   rí   s              r   Ú_prepare_qwen3_asr_audio_inputsr    sý   € ð KÐJÐJÐJÐJÐJå# E¨:Ñ6Ô6€HÝ)¨%Ð1AÑBÔB€NØÐ˜>Ð1Øˆà Ð!6Ô7ÐØ(;Ô(AÑ%€JÐ%Ø&¨8°a©<Ñ8€JØ&×*Ò*¨2Ñ.Ô.×1Ò1µ%´*Ñ=Ô=€LØ'×,Ò,¨Z¸ÀRÑHÔH×LÒLÐQSÐLÑTÔT×\Ò\Ð]_Ñ`Ô`×cÒcÕdiÔdnÑoÔo€MØ/Ð/°¸|È^Ð]eÑfÔf€Fˆ<ÑÐÐr1   c                ó  ‡— ‰                      d¦  «        €®t          | d¦  «        rž‰                      d¦  «        }‰                      d¦  «        }|du p|du p|j        d         |j        d         k    }|rNt          t	          j        | j        ¦  «        j        ¦  «        }ˆfd„|D ¦   «         } | j        d	i |¤Ž\  }}|‰d<   t           	                    ¦   «         D ],\  }	}
t          ˆfd„|	D ¦   «         ¦  «        r |
| ‰¦  «         Œ-dS )
uî  Inject precomputed tensors for data-dependent ops the model would otherwise hit during tracing.

    Two layers:
    - Outer LLM rope index (`get_rope_index`) â€” generic `hasattr` probe; covers Qwen-VL / GLM-4V etc.
    - Per-encoder preparer dispatched by marker kwargs present in `inputs` (e.g. `grid_thw`,
      `target_sizes`, `(input_features, feature_lens)`) â€” see `register_export_input_preparer`.
      A preparer fires only when every one of its markers is present in `inputs`.
    rÔ   NÚget_rope_indexÚ	input_idsÚattention_maskrI   c                ó*   •— i | ]}|‰v ¯|‰|         “ŒS r2   r2   )rl   rx   r¯   s     €r   r¹   z,precompute_export_inputs.<locals>.<dictcomp>,  s$   ø€ ÐLÐLÐL¨AÀÀVÀÀ˜1˜f QœiÀÀÀr1   c              3  óF   •K  — | ]}‰                      |¦  «        d uV — Œd S r+   )r'   )rl   Úmr¯   s     €r   rn   z+precompute_export_inputs.<locals>.<genexpr>3  s2   øè è € Ð:Ð:¨Qˆv�zŠz˜!‰}Œ} DÐ(Ð:Ð:Ð:Ð:Ð:Ð:r1   r2   )r'   ru   rû   rr   ÚinspectÚ	signaturer  r¢   rÇ   rw   Úall)r�   r¯   r  Ú	attn_maskÚ
is_prefillÚrope_paramsÚrope_inputsrÔ   rF   rÈ   Úpreparers    `         r   r½   r½     sB  ø€ ð ‡z‚z�.Ñ!Ô!Ð)­g°eÐ=MÑ.NÔ.NÐ)Ø—J’J˜{Ñ+Ô+ˆ	Ø—J’JÐ/Ñ0Ô0ˆ	Ø $Ð&Ðg¨)°tÐ*;Ðg¸y¼ÈqÔ?QÐU^ÔUdÐefÔUgÒ?gˆ
Øð 	2Ý�gÔ/°Ô0DÑEÔEÔPÑQÔQˆKØLÐLÐLÐL°ÐLÑLÔLˆKØ2˜eÔ2ÐAÐA°[ÐAÐA‰OˆL˜!Ø%1ˆF�>Ñ"õ 5×:Ò:Ñ<Ô<ð $ð $Ñˆ�ÝÐ:Ð:Ð:Ð:°'Ð:Ñ:Ô:Ñ:Ô:ð 	$ØˆH�U˜FÑ#Ô#Ð#øð$ð $r1   rÅ   c              #  óÊ   ‡‡‡K  — g Š| j         Št          j        ‰¦  «        Št          j        ‰¦  «        ˆˆˆfd„¦   «         }|| _         	 ‰V — ‰| _         dS # ‰| _         w xY w)zÞCapture forward call kwargs into a list (one dict per call).

    Positional args are normalised to kwargs via `inspect.signature` so the
    captured dicts can be passed directly as `kwargs=inputs` to `torch.export`.
    c                 óš  •— i } ‰	j         | i |¤Ž}|j                             ¦   «         D ]…\  }}‰	j        |         }|j        t
          j        j        k    r(|                     t          j
        |¦  «        ¦  «         ŒT|j        t
          j        j        k    rt          j
        |¦  «        ||<   Œ†‰                     |¦  «          ‰| i |¤ŽS r+   )ÚbindÚ	argumentsrw   r¢   Úkindr  Ú	ParameterÚVAR_KEYWORDÚupdateÚcopyÚdeepcopyÚVAR_POSITIONALr-   )
ÚargsÚkwargsÚcapturedÚboundrÂ   rƒ   ÚparamÚcallsr   Úsigs
          €€€r   Úwrapperz!_capture_forward.<locals>.wrapperK  sË   ø€ àˆØ�”˜$Ð) &Ð)Ð)ˆØ œ?×0Ò0Ñ2Ô2ð 	6ð 	6‰KˆD�%Ø”N 4Ô(ˆEØŒz�WÔ.Ô:Ò:Ð:Ø—’¥¤¨eÑ 4Ô 4Ñ5Ô5Ð5Ð5Ø”�wÔ0Ô?Ò?Ð?Ý!%¤¨uÑ!5Ô!5�˜‘øØ�Š�XÑÔÐØˆx˜Ð( Ð(Ð(Ð(r1   N)Úforwardr  r  Ú	functoolsÚwraps)rÅ   r%  r#  r   r$  s     @@@r   Ú_capture_forwardr)  ?  s”   øøøè è € ð €EØŒ~€HÝ
Ô
˜HÑ
%Ô
%€Cå„_�XÑÔð
)ð 
)ð 
)ð 
)ð 
)ð 
)ñ Ôð
)ð €F„Nð"Øˆˆˆà!ˆŒˆˆø˜ˆŒÐ!Ð!Ð!Ð!s   ÁA Á	A"r	   ú'dict[str, tuple[torch.nn.Module, dict]]c           
     óR  — 	 t          | ¦  «        5 } | j        di t          j        |¦  «        ¤dddœ¤Ž ddd¦  «         n# 1 swxY w Y   nZ# t          $ rM}t          dt          | ¦  «        j        › dt          | 	                    ¦   «         ¦  «        › d�¦  «        |‚d}~ww xY wt          |¦  «        dk     r5t          dt          | ¦  «        j        › dt          |¦  «        › d	�¦  «        ‚t          j        | ¦  «        |d
         ft          j        | ¦  «        |d         fdœS )u‘  Run `model.generate()` for 2 tokens and capture prefill and decode inputs.

    Reuses the full generation machinery so every architecture (decoder-only, SSM,
    encoder-decoder, multi-modal, â€¦) gets correct inputs without reimplementing the loop.

    Returns:
        `dict[str, tuple[torch.nn.Module, dict]]`:
        `{"prefill": (model, prefill_inputs), "decode": (model, decode_inputs)}`
    r   )Úmax_new_tokensÚmin_new_tokensNz$decompose_prefill_decode failed for ú. Inputs passed: z<. Make sure the inputs are compatible with model.generate().z6decompose_prefill_decode expected at least 2 calls to z;.forward() during generate(max_new_tokens=2), but captured z«. This likely means generate() bypasses the top-level forward() (e.g. delegates to an inner model), so prefill/decode decomposition is not supported for this architecture.r   rI   )ÚprefillÚdecoder2   )r)  Úgenerater  r  Ú	ExceptionÚRuntimeErrorrs   rM   rY   ÚkeysÚlen)r�   r¯   r#  Úes       r   Údecompose_prefill_decoder7  _  s¸  € ðÝ˜eÑ$Ô$ð 	X¨ØˆEŒNÐWÐW�Tœ]¨6Ñ2Ô2ÐWÀ1ÐUVÐWÐWÐWÐWÐWð	Xð 	Xð 	Xñ 	Xô 	Xð 	Xð 	Xð 	Xð 	Xð 	Xð 	Xøøøð 	Xð 	Xð 	Xð 	Xøøåð ð ð ÝðJµ4¸±;´;Ô3Gð Jð JÝ" 6§;¢;¡=¤=Ñ1Ô1ðJð Jð Jñ
ô 
ð ð		øøøøðøøøõ ˆ5�z„z�A‚~€~ÝðVÅTÈ%Á[Ä[ÔEYð Vð VÝ?BÀ5¹z¼zðVð Vð Vñ
ô 
ð 	
õ ”I˜eÑ$Ô$ e¨A¤hÐ/Ý”9˜UÑ#Ô# U¨1¤XÐ.ðð ð s:   ‚A ‘%A¶A ÁAÁA Á	AÁ
A Á
B%ÁAB Â B%)Úmulti_modal_projectorÚ	connectorÚembed_visionÚembed_audio)Úlm_headúdict[str, torch.nn.Module]c                ó>  — i }d}dD ](}|                       |¬¦  «        }|�|| ur
|||› d�<   d}Œ)|                      ¦   «         }|�	|| ur||d<   | | j        hD ]<}t          t          z   D ]*}||vr$t          ||d¦  «        �t          ||¦  «        ||<   Œ+Œ=|rd|vri S |S )u  Return `{attr_name: module}` for multi-modal submodules found on `model`.

    Uses the canonical `PreTrainedModel.get_encoder("image"/"audio")` and `get_decoder()`
    accessors for encoders and the language model. Projectors and `lm_head` are looked
    up by name on `model` and its `base_model` (e.g. `LlavaModel` under `LlavaForConditionalGeneration`).

    Only returns results when at least one modal encoder AND a language model are found â€”
    otherwise the model is not multi-modal and should be exported as a single unit.
    F)ÚimageÚaudio)ÚmodalityNÚ_encoderTÚlanguage_model)Úget_encoderÚget_decoderÚ
base_modelÚ_MULTIMODAL_PROJECTOR_NAMESÚ_MULTIMODAL_LM_HEAD_NAMESr   )r�   ÚfoundÚhas_encoderrA  ÚencoderÚdecoderÚrootrÂ   s           r   Ú_find_multimodal_submodulesrN  Š  s
  € ð )+€Eà€KØ&ð ð ˆØ×#Ò#¨XÐ#Ñ6Ô6ˆð Ð 7°%Ð#7Ð#7Ø+2ˆE�XÐ'Ð'Ð'Ñ(ØˆKøà×ÒÑ!Ô!€GØÐ˜w¨eÐ3Ð3Ø")ˆÐÑà˜Ô(Ð)ð 2ð 2ˆÝ/Õ2KÑKð 	2ð 	2ˆDØ˜5Ð Ð ¥W¨T°4¸Ñ%>Ô%>Ð%JÝ% d¨DÑ1Ô1��d‘øð	2ð ð Ð*°%Ð7Ð7Øˆ	à€Lr1   Úboolc                ó:   — t          t          | ¦  «        ¦  «        S )zTReturns `True` if the model is multi-modal with modal encoders and a language model.)rO  rN  r¤   s    r   Úis_multimodalrQ  ¯  s   € åÕ+¨EÑ2Ô2Ñ3Ô3Ð3r1   c           
     óŠ  ‡‡— t          | ¦  «        }|s%t          dt          | ¦  «        j        › d�¦  «        ‚	 t	          j        ¦   «         5 Št          j        ¦   «         5  ˆfd„|                     ¦   «         D ¦   «         Š | d	i t          j
        |¦  «        ¤Ž ddd¦  «         n# 1 swxY w Y   ddd¦  «         n# 1 swxY w Y   nZ# t          $ rM}t          dt          | ¦  «        j        › dt          |                     ¦   «         ¦  «        › d�¦  «        |‚d}~ww xY wˆfd„|                     ¦   «         D ¦   «         S )
u‹  Capture inputs to each multi-modal submodule via a single forward pass.

    Detects all known multi-modal submodules by attribute name (vision tower, projector,
    language model, lm_head, â€¦) and captures their forward kwargs during one
    `model(**inputs)` call.

    Each submodule is returned as a separate `name: (module, inputs)` entry for
    independent export. The token-merge step (e.g. `masked_scatter` for multi-modal models)
    is intentionally left outside the exported graphs â€” it is the caller's responsibility
    to assemble `inputs_embeds` from the encoder outputs before running the decoder.

    Returns:
        `dict[str, tuple[torch.nn.Module, dict]]`: One `name: (module, inputs)`
        entry per detected submodule (image/audio encoder, projector, language model, lm_head).

    Raises:
        `ValueError`: if no known multi-modal submodules are found on the model.
    z8decompose_multimodal found no multi-modal submodules on zB. Expected an image/audio encoder + language model, found neither.c                ó\   •— i | ](\  }}|‰                      t          |¦  «        ¦  «        “Œ)S r2   )r"   r)  )rl   rÂ   rÅ   r#   s      €r   r¹   z(decompose_multimodal.<locals>.<dictcomp>Ð  sC   ø€ ð  ð  ð  ÙHTÈÈf��e×)Ò)Õ*:¸6Ñ*BÔ*BÑCÔCð ð  ð  r1   Nz decompose_multimodal failed for r.  rA   c                óH   •— i | ]\  }}‰|         ¯||‰|         d          f“ŒS )rø   r2   )rl   rÂ   rÅ   Úsubmodule_inputss      €r   r¹   z(decompose_multimodal.<locals>.<dictcomp>Ù  sJ   ø€ ð ð ð áˆD�&Ø˜DÔ!ðØˆvÐ'¨Ô-¨bÔ1Ð2ðð ð r1   r2   )rN  rº   rs   rM   r    r!   rV   r¼   rw   r  r  r2  r3  rY   r4  )r�   r¯   Ú
submodulesr6  r#   rU  s       @@r   Údecompose_multimodalrW  ´  s  øø€ õ& -¨UÑ3Ô3€JØð 
ÝðPÅtÈEÁ{Ä{ÔG[ð Pð Pð Pñ
ô 
ð 	
ð
	ÝÔ!Ñ#Ô#ð 	+ u­e¬m©o¬oð 	+ð 	+ð ð  ð  ð  ØXb×XhÒXhÑXjÔXjð ñ  ô  Ðð ˆEÐ*Ð*•D”M &Ñ)Ô)Ð*Ð*Ð*ð		+ð 	+ð 	+ñ 	+ô 	+ð 	+ð 	+ð 	+ð 	+ð 	+ð 	+øøøð 	+ð 	+ð 	+ð 	+ð 	+ð 	+ð 	+ñ 	+ô 	+ð 	+ð 	+ð 	+ð 	+ð 	+ð 	+øøøð 	+ð 	+ð 	+ð 	+øøõ
 ð ð ð ÝØl­t°E©{¬{Ô/CÐlÐlÕVZÐ[a×[fÒ[fÑ[hÔ[hÑViÔViÐlÐlÐlñ
ô 
àð	øøøøðøøøð
ð ð ð à&×,Ò,Ñ.Ô.ðñ ô ð s`   ºC ÁB?Á!;B(ÂB?Â(B,	Â,B?Â/B,	Â0B?Â3C Â?CÃC ÃCÃC Ã
D"ÃADÄD"c                ó”   — t          | |¦  «        }|d         \  }}t          |¦  «        s|S t          ||¦  «        }|d         |d<   |S )u  Decompose a generative model into independently exportable `(model, forward_inputs)` pairs.

    Runs `decompose_prefill_decode` to capture prefill and decode forward kwargs from a real
    `model.generate(**inputs, max_new_tokens=2)`. If the prefill is multi-modal (per `is_multimodal`),
    further splits it into one entry per submodule (vision/audio encoder, projector, language model,
    `lm_head`) via `decompose_multimodal`.

    Args:
        model: Generative model. Must support `model.generate(**inputs)`.
        inputs: **Generate** kwargs â€” what you'd pass to `model.generate(**inputs)`.

    Returns:
        `{component_name: (submodel, forward_inputs)}`. Keys are `"prefill"` / `"decode"` for
        plain generative models and `"<modality>_encoder"` / `"multi_modal_projector"` /
        `"language_model"` / `"lm_head"` / `"decode"` for multi-modal generative models.
    r/  r0  )r7  rQ  rW  )r�   r¯   ÚstagesÚprefill_modelÚprefill_inputsÚ
componentss         r   Údecompose_for_generationr]  à  s[   € õ& & e¨VÑ4Ô4€FØ$*¨9Ô$5Ñ!€M�>å˜Ñ'Ô'ð Øˆå% m°^ÑDÔD€JØ! (Ô+€JˆxÑØÐr1   )r   r   r   r   r   r   )r   r   )r%   r   )r%   r   r8   r9   )r%   r   r>   r   )rD   r   )r   r   r/   rg   r8   r   )r{   )r   r   r|   r   )r   r   r8   r„   )r   r   r8   r   )r   r   r’   r“   r”   r•   r8   r   )r�   rž   r8   rŸ   )r�   rž   r8   r§   )r�   rž   r¯   r°   r8   r±   )r�   rÁ   rÂ   r   r8   rÃ   )rÈ   r   )r�   rÁ   r¯   rÍ   r8   r9   )rÅ   rÁ   )r�   r	   r¯   rÍ   r8   r*  )r�   r	   r8   r=  )r�   r	   r8   rO  )KÚ__doc__Ú
__future__r   r    r  Úenumr'  r  rò   Úcollections.abcr   Útypingr   Úutilsr   Úutils.import_utilsr   Ú
get_loggerrM   ÚloggerrV   Úmodeling_utilsr	   Úvision_utilsr
   r   r   r   r   r   r   Ú__annotations__r   r   Úcontextmanagerr   r$   r(   r4   r7   r=   rG   rC   re   rs   rf   ÚEnumÚSymIntÚSymFloatÚSymBoolrk   r€   r†   r‘   rœ   r¥   r©   r»   rÀ   rÆ   rÇ   rË   rÜ   rå   rõ   r  r½   r)  r7  rG  rH  rN  rQ  rW  r]  r2   r1   r   ú<module>ro     s  ðð ð ð ð* #Ð "Ð "Ð "Ð "Ð "à Ð Ð Ð Ø €€€Ø €€€Ø Ð Ð Ð Ø €€€Ø 
€
€
€
Ø *Ð *Ð *Ð *Ð *Ð *Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø 3Ð 3Ð 3Ð 3Ð 3Ð 3ð 
ˆÔ	˜HÑ	%Ô	%€ð ÐÑÔð Ø€L€L€Là0Ð0Ð0Ð0Ð0Ð0ðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð" 8:€Ð 9Ð 9Ð 9Ñ 9Ø,.€Ð .Ð .Ð .Ñ .Ø/1Ð Ð 1Ð 1Ð 1Ñ 1ð Ôð*ð *ð *ñ Ôð*ð Ôð	ð 	ð 	ñ Ôð	ð Ôðð ð ñ Ôððð ð ð ðð ð ð ðð ð ð ðð ð ð ð6ð ð ð ð&ð ð ð ðJ '+ WÐ Ð ,Ð ,Ð ,Ñ ,ØÐÑÔð QØ˜œ E¤L°%´.À%Ä-ÐPÑPÐðð ð ð ð09ð 9ð 9ð 9ð 9ð,
)ð 
)ð 
)ð 
)ð*ð *ð *ð *ð&)ð )ð )ð )ð	ð 	ð 	ð 	ðð ð ð ð i€ð4'ð 4'ð 4'ð 4'ð|ð ð ð ð <>Ð Ð =Ð =Ð =Ñ =ð
ð 
ð 
ð 
ð  Ð 
Ñ+Ô+ð!
ð !
ð !
ñ ,Ô+ð!
ðH  Ð Ñ/Ô/ð[ð [ð [ñ 0Ô/ð[ð$  ÐÐ 0°.ÑAÔAðSð Sð Sñ BÔAðSð>  ÐÐ 0Ð2GÑHÔHðgð gð gñ IÔHðgð*$ð $ð $ð $ðH Ôð"ð "ð "ñ Ôð"ð>"ð "ð "ð "ðN dÐ Ø(Ð ð"ð "ð "ð "ðJ4ð 4ð 4ð 4ð
)ð )ð )ð )ðXð ð ð ð ð r1   