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

多库生成PyTorch数据集的适配设计模式选型咨询

解决方案:适配多库的PyTorch数据集构建流程

针对你的需求,推荐采用策略模式 + 统一工厂的组合方案,既解决不同库的输入/流程差异问题,又能标准化最终的PyTorch Dataset输出,同时兼顾扩展性。

核心思路

我们的核心目标是:无论使用哪个第三方库,最终输出都是符合PyTorch要求的Dataset子类实例。因此可以把流程拆分为两层:

  1. 数据加载策略层:为每个库单独实现适配逻辑,处理其特定的输入要求和数据提取流程,最终输出标准化的样本-标签集合
  2. 统一工厂层:根据指定的库类型选择对应策略,用标准化数据生成最终的PyTorch Dataset

具体实现示例

1. 定义统一的PyTorch Dataset类

这是我们唯一的"产品",所有流程最终都要生成这个类的实例:

import torch
from torch.utils.data import Dataset

class CustomTorchDataset(Dataset):
    def __init__(self, samples, labels):
        self.samples = torch.tensor(samples, dtype=torch.float32)
        self.labels = torch.tensor(labels, dtype=torch.long)
    
    def __len__(self):
        return len(self.samples)
    
    def __getitem__(self, idx):
        return self.samples[idx], self.labels[idx]

2. 为每个库实现数据加载策略

用抽象基类约束策略接口,每个策略只负责对应库的逻辑:

from abc import ABC, abstractmethod

class BaseDataLoader(ABC):
    @abstractmethod
    def load_and_process(self):
        """返回标准化的(samples, labels)元组,格式统一为numpy数组/张量"""
        pass

# LibraryA的适配策略
class LibADataLoader(BaseDataLoader):
    def __init__(self, data_path, preprocess_config):
        # 接收LibraryA要求的输入参数
        self.data_path = data_path
        self.config = preprocess_config
    
    def load_and_process(self):
        # 调用LibraryA的API加载数据
        raw_data = libraryA.load_dataset(self.data_path, self.config)
        # 提取并标准化样本和标签
        samples = raw_data.features.to_numpy()
        labels = raw_data.targets.to_numpy()
        return samples, labels

# LibraryB的适配策略
class LibBDataLoader(BaseDataLoader):
    def __init__(self, api_session, query_filters):
        # 接收LibraryB要求的输入参数(和A完全不同)
        self.session = api_session
        self.filters = query_filters
    
    def load_and_process(self):
        # 调用LibraryB的API加载数据
        raw_data = libraryB.query_data(self.session, self.filters)
        # 提取并标准化样本和标签,输出格式和LibA一致
        samples = raw_data.get_feature_matrix()
        labels = raw_data.get_label_vector()
        return samples, labels

3. 统一工厂类

负责选择策略并生成最终Dataset,通过关键字参数兼容不同库的输入差异:

class DatasetFactory:
    # 用字典映射替代if-elif,扩展性更好
    _loader_registry = {
        "libA": LibADataLoader,
        "libB": LibBDataLoader
    }

    @classmethod
    def create_dataset(cls, loader_type, **kwargs):
        if loader_type not in cls._loader_registry:
            raise ValueError(f"不支持的数据加载类型: {loader_type}")
        
        # 实例化对应策略,传入该库所需的参数
        loader = cls._loader_registry[loader_type](**kwargs)
        # 获取标准化数据并生成PyTorch Dataset
        samples, labels = loader.load_and_process()
        return CustomTorchDataset(samples, labels)

4. 使用方式

# 用LibraryA生成Dataset
dataset_a = DatasetFactory.create_dataset(
    loader_type="libA",
    data_path="./data/liba_data.csv",
    preprocess_config={"normalize": True}
)

# 用LibraryB生成Dataset,输入参数完全不同
dataset_b = DatasetFactory.create_dataset(
    loader_type="libB",
    api_session=libraryB.connect("api_key_xxxx"),
    query_filters={"sensor_id": "temp_001", "time_range": ["2024-01-01", "2024-01-31"]}
)

方案优势

  • 解耦库依赖:每个加载策略只关注对应库的逻辑,工厂类只负责流程调度,避免代码冗余和强耦合
  • 灵活兼容输入:通过关键字参数接收不同库的输入要求,不需要强行统一方法签名
  • 扩展性强:新增LibraryC时,只需实现LibCDataLoader并注册到工厂的注册表,无需修改现有代码
  • 符合单一职责:数据加载、标准化、Dataset生成三个环节各自独立,维护成本低

对比其他选项

如果直接用LibADatasetFactory、LibBDatasetFactory这种强耦合的具体工厂,会导致每个工厂都要重复实现Dataset生成逻辑,代码冗余且扩展性差。而抽象工厂/工厂方法确实不适用——它们的设计意图是生成多组相关产品,而这里只有Dataset一种产品,且无法统一输入签名。

内容的提问来源于stack exchange,提问作者J. Lo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 12:51:56