迭代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
相关产品推荐
相关产品推荐

