U
    Ôd¨ip  ã                   @   s*   d dl mZ ddlmZ edœdd„ZdS )é    )Ú
DataLoaderé   )Úis_torch_xla_available)Ú
dataloaderc                 C   sd   t ƒ r\dd lm  m} t| |jƒs,tdƒ‚dd lm  m} | 	| 
¡ d¡}|| jd< | S | S d S )Nr   zPThe dataloader must be a `torch_xla.distributed.parallel_loader.MpDeviceLoader`.)ZfsdpNZinput_sharding)r   Z%torch_xla.distributed.parallel_loaderZdistributedZparallel_loaderÚ
isinstanceZMpDeviceLoaderÚAssertionErrorZtorch_xla.distributed.spmdZspmdZShardingSpecZget_global_meshZ_parallel_loader_kwargs)r   ÚplÚxsZsharding_spec© r
   úA/tmp/pip-unpacked-wheel-bm_b0l5e/transformers/integrations/tpu.pyÚtpu_spmd_dataloader   s     ÿþ
r   N)Ztorch.utils.datar   Úutilsr   r   r
   r
   r
   r   Ú<module>   s   