
    ij                     j    d dl Zd dlmZ d dlmZ d dlmZ d dlm	Z	 d dlm
Z
 d dlmZ  G d d	e      Zy)
    N)tree)Layer)SeedGenerator)backend_utils)	jax_utils)trackingc                   X     e Zd ZdZ fdZ fdZej                  dd       Zd Z	 xZ
S )	DataLayeraj
  Layer designed for safe use in `tf.data` or `grain` pipeline.

    This layer overrides the `__call__` method to ensure that the correct
    backend is used and that computation is performed on the CPU.

    The `call()` method in subclasses should use `self.backend` ops. If
    randomness is needed, define both `seed` and `generator` in `__init__` and
    retrieve the running seed using `self._get_seed_generator()`. If the layer
    has weights in `__init__` or `build()`, use `convert_weight()` to ensure
    they are in the correct backend.

    **Note:** This layer and its subclasses only support a single input tensor.

    Examples:

    **Custom `DataLayer` subclass:**

    ```python
    from keras.src.layers.preprocessing.data_layer import DataLayer
    from keras.src.random import SeedGenerator


    class BiasedRandomRGBToHSVLayer(DataLayer):
        def __init__(self, seed=None, **kwargs):
            super().__init__(**kwargs)
            self.probability_bias = ops.convert_to_tensor(0.01)
            self.seed = seed
            self.generator = SeedGenerator(seed)

        def call(self, inputs):
            images_shape = self.backend.shape(inputs)
            batch_size = 1 if len(images_shape) == 3 else images_shape[0]
            seed = self._get_seed_generator(self.backend._backend)

            probability = self.backend.random.uniform(
                shape=(batch_size,),
                minval=0.0,
                maxval=1.0,
                seed=seed,
            )
            probability = self.backend.numpy.add(
                probability, self.convert_weight(self.probability_bias)
            )
            hsv_images = self.backend.image.rgb_to_hsv(inputs)
            return self.backend.numpy.where(
                probability[:, None, None, None] > 0.5,
                hsv_images,
                inputs,
            )

        def compute_output_shape(self, input_shape):
            return input_shape
    ```

    **Using as a regular Keras layer:**

    ```python
    import numpy as np

    x = np.random.uniform(size=(1, 16, 16, 3)).astype("float32")
    print(BiasedRandomRGBToHSVLayer()(x).shape)  # (1, 16, 16, 3)
    ```

    **Using in a `tf.data` pipeline:**

    ```python
    import tensorflow as tf

    tf_ds = tf.data.Dataset.from_tensors(x)
    tf_ds = tf_ds.map(BiasedRandomRGBToHSVLayer())
    print([x.shape for x in tf_ds])  # [(1, 16, 16, 3)]
    ```

    **Using in a `grain` pipeline:**

    ```python
    import grain

    grain_ds = grain.MapDataset.source([x])
    grain_ds = grain_ds.map(BiasedRandomRGBToHSVLayer())
    print([x.shape for x in grain_ds])  # [(1, 16, 16, 3)]
    c                 d    t        |   di | t        j                         | _        d| _        y )NT )super__init__r   DynamicBackendbackend!_allow_non_tensor_positional_args)selfkwargs	__class__s     ~/var/www/html/emotional.easysim.app/public_html/venv/lib/python3.12/site-packages/keras/src/layers/preprocessing/data_layer.pyr   zDataLayer.__init__^   s+    "6"$33515.    c                 *    t        j                  |      d   }t        |t        j                        st        j                         rt        j                  |      s j                  j                  d       t        j                   fd|      }d} j                  r	d _        d}	 t         8  |fi |} j                  j                          |rd _        |S t        |t        j                        sWt        j                          rCt        j"                  j                  j%                  d      5  t         8  |fi |cd d d        S t         8  |fi |S #  j                  j                          |rd _        w w xY w# 1 sw Y   y xY w)Nr   
tensorflowc                 R    j                   j                  | j                        S )N)dtype)r   convert_to_tensorcompute_dtype)xr   s    r   <lambda>z$DataLayer.__call__.<locals>.<lambda>m   s&    $,,88T// 9  r   FTcpu)r   flatten
isinstancekerasKerasTensorr   in_tf_graphr   is_in_jax_tracing_scoper   set_backendmap_structure_convert_input_argsr   __call__resetin_grain_data_pipelinesrcdevice_scope)r   inputsr   sample_inputswitch_convert_input_argsoutputsr   s   `     r   r)   zDataLayer.__call__c   sc   ||F+A.<):):;))+55lC LL$$\2'' 	F ).%''+0(,0)4'*6<V<""$,/3D,N<):):;446 ""//6 :w'9&9: : 7#F5f55 ""$,/3D, -: :s   *E  6F	 &F	Fc                 j   t        | d      rt        | d      st        d      |!|t        j                  j                         k(  r| j                  S t        | d      si | _        || j
                  v r| j
                  |   S t        | j                  | j                        }|| j
                  |<   |S )Nseed	generatorzpThe `seed` and `generator` variable must be set in the `__init__` method before calling `_get_seed_generator()`._backend_generators)r   )hasattr
ValueErrorr"   r   r4   r5   r   r3   )r   r   seed_generators      r   _get_seed_generatorzDataLayer._get_seed_generator   s    tV$GD+,FL  ?g)>)>)@@>>!t23')D$d...++G44&tyy$,,G,:  )r   c                     | j                   j                  t        j                   j                         k(  r|S t        j                  j	                  |      }| j                   j                  |      S )z9Convert the weight if it is from the a different backend.)r   namer"   opsconvert_to_numpyr   )r   weights     r   convert_weightzDataLayer.convert_weight   sO    << 5 5 77MYY//7F<<11&99r   )N)__name__
__module____qualname____doc__r   r)   r    no_automatic_dependency_trackingr9   r?   __classcell__)r   s   @r   r
   r
   
   s4    Qf6
#6J .. / :r   r
   )keras.src.backendr"   	keras.srcr   keras.src.layers.layerr   keras.src.random.seed_generatorr   keras.src.utilsr   r   r   r
   r   r   r   <module>rK      s(      ( 9 ) % $U: U:r   