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

