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

PyTorch GAN生成器用GloVe预训练Embedding层维度不匹配如何解决

问题根因

你遇到的报错核心是张量维度不匹配,同时生成器的架构、嵌入层的放置逻辑也存在问题:

  • 嵌入层放在生成器最前端,输入的[53, 20]词索引会被转换为[53, 20, 25]的3维嵌入向量,后续线性层默认作用于最后一维,最终输出仍为3维;而真实样本是2维的词索引张量,维度数不一致无法拼接。
  • 生成器的输入使用全零张量不符合GAN的潜空间采样逻辑,且线性层的维度设置把批次大小作为固定参数,会导致后续批次大小变化时模型不可用。
修复方案

这里提供两种常用的文本GAN实现思路,你可以根据需求选择:

方案1:嵌入层移到判别器侧(实现最简单,无需处理离散梯度问题)

这种方案生成器直接输出和真实样本形状完全一致的词索引,嵌入层统一放在判别器的最前端,对所有输入做嵌入转换:

import torch
import torch.nn as nn

# 固定参数定义
vocab_size = 5119
seq_len = 53
batch_size = 20
emb_dim = 25
latent_dim = 100 # 潜空间维度可自行调整

# 生成器定义
class Generator(nn.Module):
    def __init__(self):
        super().__init__()
        self.model = nn.Sequential(
            nn.Linear(latent_dim, 256),
            nn.ReLU(),
            nn.Linear(256, seq_len * vocab_size)
        )
        
    def forward(self, x):
        # x为随机采样的潜空间噪声,形状[batch_size, latent_dim]
        output = self.model(x)
        output = output.view(batch_size, seq_len, vocab_size)
        # 取概率最高的词索引,调整维度和真实样本一致为[seq_len, batch_size]
        output = output.argmax(dim=-1).permute(1, 0)
        return output

# 判别器定义
emb = nn.Embedding.from_pretrained(torch.FloatTensor(weights_matrix))
class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.emb = emb
        self.classifier = nn.Sequential(
            nn.Linear(emb_dim, 64),
            nn.ReLU(),
            nn.Linear(64, 1),
            nn.Sigmoid()
        )
        
    def forward(self, x):
        # 输入x形状为[seq_len, batch_size],和真实样本格式一致
        x_emb = self.emb(x) # 转换为[seq_len, batch_size, emb_dim]
        pred = self.classifier(x_emb)
        # 按序列维度取平均,得到每个样本的判别结果[batch_size, 1]
        return pred.mean(dim=0)

使用时生成样本和真实样本形状完全一致,可直接拼接:

# 潜空间采样
latent_samples = torch.randn((batch_size, latent_dim))
generated_samples = generator(latent_samples)
# 拼接,形状为[106, 20]
all_samples = torch.cat((real_samples, generated_samples))

方案2:生成器侧将嵌入映射回词索引(支持端到端训练)

如果你需要保留嵌入层在生成器内,需要将最终输出的25维嵌入映射回词汇表索引,由于argmax是离散不可导操作,需要用Gumbel-Softmax做可微近似,最终输出调整为[53,20]的2维张量即可和真实样本拼接。

内容的提问来源于stack exchange,提问作者Evi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 08:06:07