o
    i e                     @   s<   d Z ddlZddlm  mZ ddlmZ G dd dZ	dS )z:Utility object to handler partial batches for TPUStrategy.    N)backendc                   @   s8   e Zd ZdZdd Zdd Zdd Zdd	 Zd
d ZdS )PartialBatchPaddingHandlerzBA container that holds info about partial batches for `predict()`.c                 C   s   d| _ td| _|| _d S )Nr   )padded_batch_sizetfzerospadding_maskoutput_shape)selfr    r
   ]/var/www/myenv/lib/python3.10/site-packages/keras/src/engine/partial_batch_padding_handler.py__init__   s   
z#PartialBatchPaddingHandler.__init__c                 C   sJ   t |ttfr|d }tj|sJ dd }tjt||d ddS )z>Returns the number of elements in a potentially partial batch.r   c                 S   s*   dd t j| D }|std|d S )Nc                 S   s   g | ]	}t |r|qS r
   )r   	is_tensor).0xr
   r
   r   
<listcomp>'   s
    
z\PartialBatchPaddingHandler.get_real_batch_size.<locals>._find_any_tensor.<locals>.<listcomp>z(Cannot find any Tensor in features dict.r   )r   nestflatten
ValueError)batch_featurestensorsr
   r
   r   _find_any_tensor&   s   
zHPartialBatchPaddingHandler.get_real_batch_size.<locals>._find_any_tensorint64)dtype)	
isinstancetuplelistr   r   r   r   castshape)r	   dataset_batchr   r
   r
   r   get_real_batch_size   s   z.PartialBatchPaddingHandler.get_real_batch_sizec                 C   sD   |  |}| j| }tjt|t|gdd}tj||gddS )z?Calculate and cache the amount of padding required for a batch.r   axis)r   r   r   concatenater   onesr   )r	   r   r   original_batch_sizemissing_countmaskr
   r
   r   update_mask2   s   

z&PartialBatchPaddingHandler.update_maskc                    sJ    fdd t |dkr |d S g }|D ]	}| | qt|S )z@Pads the batch dimension of a tensor to the complete batch size.c                    s   i }t | tr|  D ]
\}} |||< q|S t| j}|dks#J j|  }td|ggddgg|d   }t	
| |dS )z>Helper function to pad nested data within each batch elements.r      constant)r   dictitemslenr   r   r   r   stackr   pad)batchpadded_dict_batchkeyvaluerankr%   padding_padr	   r
   r   r6   >   s   

z2PartialBatchPaddingHandler.pad_batch.<locals>._padr(   r   )r,   appendr   )r	   dataset_batch_elementsbatch_elementsbatch_elementr
   r5   r   	pad_batch;   s   z$PartialBatchPaddingHandler.pad_batchc              	   C   s   t | j}t|jdksJ t| jdkr7tj|t|dt| dd}|jd dkr5tj	|dd}|S g }t
t| jD ]}|| }tj|t|dt| dd}|t	| q@|S )z;Removes prediction output that corresponds to padded input.r(   Nr   r    )r   	get_valuer   r,   r   r   nptakenonzerosqueezeranger7   )r	   prediction_resultr   
predictionpredictionsir
   r
   r   
apply_maskX   s*   z%PartialBatchPaddingHandler.apply_maskN)	__name__
__module____qualname____doc__r   r   r'   r;   rF   r
   r
   r
   r   r      s    	r   )
rJ   numpyr=   tensorflow.compat.v2compatv2r   	keras.srcr   r   r
   r
   r
   r   <module>   s
   