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

PyTorch中拼接不同长度1D ECG信号生成张量适配AdaptiveAveragePooling1d的方法

变长ECG信号输入AdaptiveAveragePooling1d的无填充解决方案

报错原因

直接调用torch.Tensor([array1, array2, array3])构造张量时,PyTorch会默认生成维度规整的高维张量,由于三个ECG样本的时间维度长度不一致,无法自动对齐,因此触发ValueError维度不匹配报错。

解决方案(无需填充/插值,无需调整QRS标签)

核心逻辑是不强行对齐原始变长信号的长度,利用AdaptiveAveragePooling1d支持任意长度输入、输出固定长度特征的特性,先对批次内单样本单独做池化得到固定长度特征,再堆叠为规整批次送入后续网络层。

1. 自定义数据集与collate函数

from torch.utils.data import Dataset, DataLoader
import torch
import numpy as np

class ECGDataset(Dataset):
    def __init__(self, ecg_arrays, qrs_labels):
        # 单个ECG转为(通道数, 时间步)格式,符合Pool1d输入要求
        self.ecg_list = [torch.tensor(arr, dtype=torch.float32).permute(1,0) for arr in ecg_arrays]
        self.labels = torch.tensor(qrs_labels, dtype=torch.float32)
    
    def __len__(self):
        return len(self.ecg_list)
    
    def __getitem__(self, idx):
        return self.ecg_list[idx], self.labels[idx]

# 自定义collate_fn,返回变长样本列表与对齐的标签张量
def custom_collate(batch):
    ecg_samples = [item[0] for item in batch]
    qrs_labels = torch.stack([item[1] for item in batch])
    return ecg_samples, qrs_labels

2. 自定义模型前向逻辑

class ECGQRSModel(torch.nn.Module):
    def __init__(self, pool_output_len=200, hidden_dim=128):
        super().__init__()
        self.adaptive_pool = torch.nn.AdaptiveAvgPool1d(output_size=pool_output_len)
        # 后续全连接层,输入维度为通道数*池化输出长度
        self.fc1 = torch.nn.Linear(1 * pool_output_len, hidden_dim)
        self.relu = torch.nn.ReLU()
        self.fc2 = torch.nn.Linear(hidden_dim, 1) # 回归输出QRS区间长度(毫秒)
    
    def forward(self, batch_ecgs):
        # batch_ecgs为长度等于batch_size的列表,每个元素是(1, 任意时间步)的ECG张量
        pooled_features = []
        for single_ecg in batch_ecgs:
            # 单样本过自适应池化,得到固定长度特征(1, pool_output_len)
            pooled = self.adaptive_pool(single_ecg)
            pooled_features.append(pooled.flatten())
        # 所有特征堆叠为规整批次张量:(batch_size, 1*pool_output_len)
        batch_features = torch.stack(pooled_features)
        # 送入后续层计算
        x = self.relu(self.fc1(batch_features))
        return self.fc2(x)

3. 训练调用示例

# 示例输入与标签
array1 = np.random.randn(1200,1)
array2 = np.random.randn(950,1)
array3 = np.random.randn(1000,1)
ecg_dataset = [array1, array2, array3]
# 示例QRS标签,单位毫秒
qrs_labels = [82, 95, 78]

# 构造数据集与加载器
dataset = ECGDataset(ecg_dataset, qrs_labels)
dataloader = DataLoader(dataset, batch_size=2, collate_fn=custom_collate, shuffle=True)

# 初始化模型与训练流程
model = ECGQRSModel()
loss_fn = torch.nn.MSELoss() # 回归任务用均方误差损失
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

for epoch in range(10):
    for ecgs, labels in dataloader:
        optimizer.zero_grad()
        pred = model(ecgs)
        loss = loss_fn(pred.squeeze(), labels)
        loss.backward()
        optimizer.step()

方案说明

  • 全程未对原始ECG信号做零值填充、插值拉伸操作,不会引入人工噪声,也不需要调整原始QRS标签的数值,完全保留标签的物理意义。
  • 批次规模在32以下时,逐样本池化的运算开销可忽略,若需要更高并行效率,可通过CUDA加速进一步提升运算速度。

内容的提问来源于stack exchange,提问作者Sara De Luca

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 19:36:01