You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

迭代PyTorch DataLoader时出现urllib3协议错误求助

PyTorch DataLoader多进程读取S3数据时的urllib3协议错误解决方案

问题背景

使用PyTorch DataLoader(num_workers>0多进程模式)结合smart_open读取AWS S3中的图像数据时,无规律触发urllib3.exceptions.ProtocolError,底层错误为ConnectionResetError(104, 'Connection reset by peer'),错误随机出现,难以定位。

相关代码片段

def __getitem__(self, index):
        
      pair_key = self.list_files[index]
      pair = self.s3_client.list_objects(Bucket=self.bucket_name, Prefix=pair_key, Delimiter='/')

      input_image_key = pair.get('Contents')[1].get('Key')
      input_image_path = f's3://{self.bucket_name}/{input_image_key}'
      input_image_s3_source = get_file_from_filepath(input_image_path)
      pil_input_image = Image.open(input_image_s3_source)

      target_image_key = pair.get('Contents')[0].get('Key')
      target_image_path = f's3://{self.bucket_name}/{target_image_key}'
      target_image_s3_source = get_file_from_filepath(target_image_path)
      pil_target_image = Image.open(target_image_s3_source)
      
      input_image = self.transform(pil_input_image)
      target_image = self.transform(pil_target_image)

      return input_image, target_image

def main():
  train_loader = DataLoader(
      train_dataset,
      batch_size=args.batch_size,
      shuffle=False,
      num_workers=args.num_workers,
      drop_last=True,
      pin_memory=True
  )
  print("Length train loader ",len(train_loader))
  val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, pin_memory=True, drop_last = True)

def train_fn():
    loop = tqdm(loader, leave=True)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    for idx, (x, y) in enumerate(loop):
        ...

错误栈信息

Traceback (most recent call last):
  File "main.py", line 623, in <module>
    g_scaler=g_scaler, d_scaler=d_scaler, runtime_log_folder=runtime_log_folder, runtime_log_file_name=runtime_log_file_name)
  File "main.py", line 223, in train_fn
    for idx, (x, y) in enumerate(loop):
  File "/opt/conda/lib/python3.6/site-packages/tqdm/std.py", line 1171, in __iter__
    for obj in iterable:
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/dataloader.py", line 525, in __next__
    (data, worker_id) = self._next_data()
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/dataloader.py", line 1252, in _next_data
    return (self._process_data(data), w_id)
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/dataloader.py", line 1299, in _process_data
    data.reraise()
  File "/opt/conda/lib/python3.6/site-packages/torch/_utils.py", line 429, in reraise
    raise self.exc_type(msg)

urllib3.exceptions.ProtocolError: Caught ProtocolError in DataLoader worker process 3.

Original Traceback (most recent call last):
  File "/opt/conda/lib/python3.6/site-packages/urllib3/response.py", line 436, in _error_catcher
    yield
  File "/opt/conda/lib/python3.6/site-packages/urllib3/response.py", line 518, in read
    data = self._fp.read(amt) if not fp_closed else b""
  File "/opt/conda/lib/python3.6/http/client.py", line 463, in read
    n = self.readinto(b)
  File "/opt/conda/lib/python3.6/http/client.py", line 507, in readinto
    n = self.fp.readinto(b)
  File "/opt/conda/lib/python3.6/socket.py", line 586, in readinto
    return self._sock.recv_into(b)
  File "/opt/conda/lib/python3.6/ssl.py", line 1012, in recv_into
    return self.read(nbytes, buffer)
  File "/opt/conda/lib/python3.6/ssl.py", line 874, in read
    return self._sslobj.read(len, buffer)
  File "/opt/conda/lib/python3.6/ssl.py", line 631, in read
    v = self._sslobj.read(len, buffer)

ConnectionResetError: [Errno 104] Connection reset by peer

During handling of the above exception, another exception occurred:

Traceback (most recent call last):
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/_utils/worker.py", line 210, in _worker_loop
    data = fetcher.fetch(index)
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/_utils/fetch.py", line 44, in fetch
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/opt/conda/lib/python3.6/site-packages/torch/utils/data/_utils/fetch.py", line 44, in <listcomp>
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/opt/ml/code/ImageDataset.py", line 137, in __getitem__
    pil_target_image = Image.open(target_image_s3_source)
  File "/opt/conda/lib/python3.6/site-packages/PIL/Image.py", line 2984, in open
    prefix = fp.read(16)
  File "/opt/conda/lib/python3.6/site-packages/smart_open/s3.py", line 511, in read
    self._fill_buffer(size)
  File "/opt/conda/lib/python3.6/site-packages/smart_open/s3.py", line 622, in _fill_buffer
    bytes_read = self._buffer.fill(self._raw_reader)
  File "/opt/conda/lib/python3.6/site-packages/smart_open/bytebuffer.py", line 152, in fill
    new_bytes = source.read(size)
  File "/opt/conda/lib/python3.6/site-packages/smart_open/s3.py", line 423, in read
    binary = self._read_from_body(size)
  File "/opt/conda/lib/python3.6/site-packages/smart_open/s3.py", line 411, in _read_from_body
    binary = self._body.read(size)
  File "/opt/conda/lib/python3.6/site-packages/botocore/response.py", line 77, in read
    chunk = self._raw_stream.read(amt)
  File "/opt/conda/lib/python3.6/site-packages/urllib3/response.py", line 540, in read
    raise IncompleteRead(self._fp_bytes_read, self.length_remaining)
  File "/opt/conda/lib/python3.6/contextlib.py", line 99, in __exit__
    self.gen.throw(type, value, traceback)
  File "/opt/conda/lib/python3.6/site-packages/urllib3/response.py", line 454, in _error_catcher
    raise ProtocolError("Connection broken: %r" % e, e)

urllib3.exceptions.ProtocolError: ("Connection broken: ConnectionResetError(104, 'Connection reset by peer')", ConnectionResetError(104, 'Connection reset by peer'))

解决方案

1. 调整DataLoader进程参数

多进程下大量并发S3请求容易触发连接重置,限制worker数量并改用spawn模式避免连接继承冲突:

train_loader = DataLoader(
    train_dataset,
    batch_size=args.batch_size,
    shuffle=False,
    num_workers=min(args.num_workers, 4),  # 限制worker数量,避免并发请求过多
    drop_last=True,
    pin_memory=True,
    multiprocessing_context='spawn'  # 使用spawn模式,避免fork导致的连接共享问题
)

2. 为每个worker创建独立S3客户端

全局共享S3客户端会导致进程间连接冲突,每个worker单独初始化客户端:

from botocore.config import Config

def get_s3_client():
    # 配置基础重试和超时,提升连接稳定性
    config = Config(
        retries={'max_attempts': 5, 'mode': 'standard'},
        connect_timeout=10,
        read_timeout=30
    )
    return boto3.client('s3', config=config)

def __getitem__(self, index):
    # 每个worker进程单独创建客户端
    s3_client = get_s3_client()
    pair_key = self.list_files[index]
    pair = s3_client.list_objects(Bucket=self.bucket_name, Prefix=pair_key, Delimiter='/')
    
    # 后续图像读取逻辑不变
    input_image_key = pair.get('Contents')[1].get('Key')
    input_image_path = f's3://{self.bucket_name}/{input_image_key}'
    input_image_s3_source = get_file_from_filepath(input_image_path)
    pil_input_image = Image.open(input_image_s3_source)
    
    # ... 其余代码

3. 给smart_open添加重试和超时配置

优化smart_open的S3读取参数,增强容错性:

from smart_open import open
from botocore.config import Config

def get_file_from_filepath(filepath):
    return open(
        filepath,
        'rb',
        transport_params={
            'client_kwargs': {
                'config': Config(
                    retries={'max_attempts': 5},
                    connect_timeout=10,
                    read_timeout=30
                )
            }
        }
    )

4. 预加载数据到本地缓存

如果网络环境不稳定,将S3数据提前下载到本地缓存,彻底规避网络连接问题:

import os

def __init__(self, bucket_name, list_files, transform, cache_dir='/tmp/s3_image_cache'):
    self.bucket_name = bucket_name
    self.list_files = list_files
    self.transform = transform
    self.cache_dir = cache_dir
    os.makedirs(cache_dir, exist_ok=True)
    self.s3_client = get_s3_client()  # 使用带重试的客户端

def _download_to_cache(self, s3_key):
    # 替换路径中的斜杠,避免本地目录结构混乱
    local_filename = s3_key.replace('/', '_')
    local_path = os.path.join(self.cache_dir, local_filename)
    if not os.path.exists(local_path):
        self.s3_client.download_file(self.bucket_name, s3_key, local_path)
    return local_path

def __getitem__(self, index):
    pair_key = self.list_files[index]
    pair = self.s3_client.list_objects(Bucket=self.bucket_name, Prefix=pair_key, Delimiter='/')
    
    input_image_key = pair.get('Contents')[1].get('Key')
    local_input_path = self._download_to_cache(input_image_key)
    pil_input_image = Image.open(local_input_path)
    
    target_image_key = pair.get('Contents')[0].get('Key')
    local_target_path = self._download_to_cache(target_image_key)
    pil_target_image = Image.open(local_target_path)
    
    # ... 后续图像转换逻辑

内容的提问来源于stack exchange,提问作者CoderMath

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 07:36:20