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

自定义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设置,训练时仍出现形状不匹配错误。请问是否存在我忽略的不一致来源?有哪些方法能确保真实与虚假批次始终对齐?


问题分析与解决建议

可能的不一致来源

  1. 生成器输入未同步切片:_collate_fn中仅对判别器相关张量做了min_len切片,但生成器部分的known、unknown、gen_in_fakez等完全未处理,若采样函数返回的长度与判别器张量不一致,会直接导致真实/虚假数据形状不匹配。
  2. min_len计算逻辑冗余:将固定长度的disc_in_perm(由torch.randperm(batch_size)生成,形状固定为[batch_size])加入min_len计算,可能掩盖其他数据张量的真实长度问题,导致切片后的长度不符合实际需求。
  3. DataLoader与_collate_fn的batch_size不一致:若_prepare_training中创建DataLoader时使用的batch_size与train函数传入的参数存在偏差(如分布式训练的批次拆分、梯度累积处理),会导致_collate_fn中len(batch)不等于预期值,进而引发采样张量长度异常。
  4. run_step切片逻辑不统一:若run_step中对真实/虚假数据的切片长度、维度不统一,仅处理了部分张量,也会造成形状不匹配。

解决方法

  1. 同步生成器与判别器的切片处理:在_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]
    
  2. 修正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}"
    
  3. 确保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}"
    
  4. 统一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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 03:05:59