自定义CTGAN模型训练中真实与虚假张量形状不匹配问题排查
问题描述
我正在开发适用于表格数据的自定义CTGAN类模型,实现了专属的_collate_fn与run_step()逻辑,用于控制采样、噪声注入及判别器输入。但训练过程中,尤其是判别器步骤,始终出现真实与虚假数据批次的形状不匹配错误。
已采取的解决措施:
- 在DataLoader中设置
drop_last=True; - 在
_collate_fn中将所有张量(包括disc_in_known、disc_in_unknown、disc_in_fakez、disc_in_c、disc_in_perm)切片至统一的min_len; - 在
run_step()中计算损失前对真实/虚假张量进行切片。
但仍出现如下错误:
ValueError: shape mismatch: real_mean=torch.Size([6823]), fake_mean=torch.Size([6814])
使用自定义collate_fn()是因为需要在每个批次中多次采样以进行判别,注入高斯噪声并将所有输入切片至相同最小长度以避免不匹配,该函数会在传入训练循环前准备好判别器与生成器的输入。相关代码如下:
def _collate_fn(self, batch: List[Tuple[Tensor, ...]]) -> Tuple[Tensor, ...]: batch_size = len(batch) mean = torch.zeros(batch_size, self._embedding_dim) std = mean + 1 disc_in_known, disc_in_unknown, disc_in_fakez = [], [], [] disc_in_c, disc_in_perm = [], [] for _ in range(self._discriminator_step): fakez = torch.normal(mean=mean, std=std) c1, m1, col, opt = self._sampler.sample_condvec(batch_size) known, unknown = self._sampler.sample_data(batch_size, col, opt) perm = torch.randperm(batch_size) disc_in_known.append(known) disc_in_unknown.append(unknown) disc_in_fakez.append(fakez) disc_in_c.append(c1) disc_in_perm.append(perm) # Use min_len to align shapes across all tensors min_len = min(k.shape[0] for k in disc_in_known + disc_in_fakez + disc_in_c + disc_in_perm + disc_in_unknown) disc_in_known = torch.stack([k[:min_len] for k in disc_in_known]) disc_in_unknown = torch.stack([k[:min_len] for k in disc_in_unknown]) disc_in_fakez = torch.stack([k[:min_len] for k in disc_in_fakez]) disc_in_c = torch.stack([k[:min_len] for k in disc_in_c]) disc_in_perm = torch.stack([k[:min_len] for k in disc_in_perm]) # generator input gen_in_fakez = torch.normal(mean=mean, std=std) c1, m1, col, opt = self._sampler.sample_condvec(batch_size) known, unknown = self._sampler.sample_data(batch_size, col, opt) return disc_in_known, disc_in_unknown, disc_in_fakez, disc_in_c, disc_in_perm, known, unknown, gen_in_fakez, c1, m1
训练函数代码:
def train(self, known: Tensor, unknown: Tensor, epochs: int = 10, batch_size: int = 100, shuffle: bool = True, save_freq: int = 100, resume: bool = True, lae_epochs: int = 10): #Toegevoegd #self._construct_sampler(known, unknown) self._calculate_corr(known, unknown) (epoch, global_step), dataloader = self._prepare_training(known, unknown, batch_size, shuffle, resume) if not self._lae_trained and self._known_dim != 0 and self._lae is not None: lae_dataloader = self._lae_dataloader(known, unknown, batch_size, shuffle) self._train_lae(lae_dataloader, lae_epochs) self._run_training(epoch, epochs, global_step, save_freq, dataloader)
尽管做了所有切片操作与drop_last=True设置,训练时仍出现形状不匹配错误。请问是否存在我忽略的不一致来源?有哪些方法能确保真实与虚假批次始终对齐?
问题分析与解决建议
可能的不一致来源
- 生成器输入未同步切片:
_collate_fn中仅对判别器相关张量做了min_len切片,但生成器部分的known、unknown、gen_in_fakez等完全未处理,若采样函数返回的长度与判别器张量不一致,会直接导致真实/虚假数据形状不匹配。 - min_len计算逻辑冗余:将固定长度的
disc_in_perm(由torch.randperm(batch_size)生成,形状固定为[batch_size])加入min_len计算,可能掩盖其他数据张量的真实长度问题,导致切片后的长度不符合实际需求。 - DataLoader与_collate_fn的batch_size不一致:若
_prepare_training中创建DataLoader时使用的batch_size与train函数传入的参数存在偏差(如分布式训练的批次拆分、梯度累积处理),会导致_collate_fn中len(batch)不等于预期值,进而引发采样张量长度异常。 - run_step切片逻辑不统一:若
run_step中对真实/虚假数据的切片长度、维度不统一,仅处理了部分张量,也会造成形状不匹配。
解决方法
- 同步生成器与判别器的切片处理:在
_collate_fn的生成器输入部分,使用相同的min_len对所有张量切片:# generator input gen_in_fakez = torch.normal(mean=mean, std=std)[:min_len] c1, m1, col, opt = self._sampler.sample_condvec(batch_size) c1 = c1[:min_len] m1 = m1[:min_len] known, unknown = self._sampler.sample_data(batch_size, col, opt) known = known[:min_len] unknown = unknown[:min_len] - 修正min_len计算逻辑:仅关注数据相关张量的长度,移除固定长度张量的干扰,并添加采样长度校验:
# 只计算数据张量的最小长度 min_len = min(k.shape[0] for k in disc_in_known + disc_in_unknown + disc_in_fakez) # 采样后校验长度,从源头避免异常 known, unknown = self._sampler.sample_data(batch_size, col, opt) assert known.shape[0] == batch_size, f"Known data length {known.shape[0]} != batch_size {batch_size}" assert unknown.shape[0] == batch_size, f"Unknown data length {unknown.shape[0]} != batch_size {batch_size}" - 确保batch_size一致性:在
_prepare_training中校验DataLoader的batch_size与传入参数一致,并在_collate_fn开头添加断言:# 在_train函数中设置预期批次大小 self._expected_batch_size = batch_size # 在_collate_fn开头添加校验 assert len(batch) == self._expected_batch_size, f"Collate batch size {len(batch)} != expected {self._expected_batch_size}" - 统一run_step切片逻辑:在
run_step中使用相同的长度参数对真实/虚假数据切片,并添加形状校验:# 计算损失前校验形状 assert real_data.shape == fake_data.shape, f"Real shape {real_data.shape} != Fake shape {fake_data.shape}"
内容的提问来源于stack exchange,提问作者Leyla Elkhamlichi
相关产品推荐
相关产品推荐

