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

Autoencoder求助:one-hot编码瑞利衰落实复转换训练报错

通信自编码器非训练信道层实现故障排查

设计目标

实现端到端通信自编码器,在无训练参数的中间层完成实复转换、瑞利衰落乘性干扰、复高斯白噪声加性干扰三个核心环节,网络结构分为三部分:

  • 编码器:输入为M×1维度的one-hot编码向量(M取16或64),经多层全连接层输出后转复值格式,归一化后送入信道层
  • 中间信道层:无可训练参数,对输入复值信号乘以复瑞利衰落系数,叠加复高斯白噪声后输出给解码器
  • 解码器:将输入复值信号转回实值格式,经多层全连接层后接softmax激活输出分类结果

故障现象

训练启动即抛出维度不匹配错误,设置batch_size=399时,输出logits维度为[114,64],与标签维度[399,64]无法广播对齐,报错核心信息如下:

InvalidArgumentError: Graph execution error:
Node: 'categorical_crossentropy/softmax_cross_entropy_with_logits'
logits and labels must be broadcastable: logits_size=[114,64] labels_size=[399,64]

原始实现代码

import keras
from keras.layers import Input, Dense, GaussianNoise,Lambda,Dropout, Concatenate
from keras.models import Model
import numpy as np
from numpy import sum, isrealobj, sqrt
from keras import regularizers
from numpy.random import standard_normal
from tensorflow.keras.layers import BatchNormalization
from tensorflow.keras.optimizers import Adam,SGD
from keras import backend as K
import matplotlib.pyplot as plt


# defining parameters
M = 64 # M= Number of messages to encode (Here taking 64 for 64QAM)
k = np.log2(M)
k = int(k)
n_channel = 7
R = k/n_channel
print ('M:',M,'   k:',k, '   n_channel:',n_channel,'   R', R)
EbNo=10.0**(15/10.0)
noise_std = np.sqrt(1/(2*R*EbNo)) #(Beta=(2*R*EbNo)^-1)


#generating data of size N
N = 40000
label = np.random.randint(M,size=N)


#creating one hot encoded vectors
data = []
for i in label:
    temp = np.zeros(M)
    temp[i] = 1
    data.append(temp)


data = np.array(data)
print (data.shape)


#To check randomly generated data
data_check = [28,1608,2730,3978,4620,7018,12359,17334,19173]
for i in data_check:
  print(label[i],data[i])


#defining real to complex and back.
def real_to_complex(x):
    real = x[:,0]
    imag = x[:,1]
    return tf.reshape(tf.dtypes.complex(real,imag),shape=[-1,7])

def complex_to_real(x):
    real = tf.math.real(x)
    imag = tf.math.imag(tf.dtypes.cast(x,tf.complex64))
    real_expand = tf.expand_dims(real,-1)
    imag_expand = tf.expand_dims(imag,-1)
    concated = tf.concat([real_expand, imag_expand],-1)
    return tf.reshape(concated,shape=[-1,7])


# Create random Complex Channel
h_real = 1/np.sqrt(2)*K.random_normal((n_channel,),mean=0,stddev=1)
h_imag = 1/np.sqrt(2)*K.random_normal((n_channel,),mean=0,stddev=1)
h = tf.dtypes.complex(h_real,h_imag)
 
 # Create random Complex Gaussian Noise
noise_real = 1/np.sqrt(2)*K.random_normal((n_channel,),mean=0,stddev=noise_std)
noise_imag = 1/np.sqrt(2)*K.random_normal((n_channel,),mean=0,stddev=noise_std)
noise = tf.dtypes.complex(noise_real,noise_imag)



# Autoencoder structure

###Encoder###
#R = k/7
#n_channel = 7
print (int(k/R))
input_data = Input(shape=(M,))
encoded = Dense(M, activation='relu')(input_data)
encoded1 = Dense(2*n_channel, activation='linear')(encoded)
encoded2 = BatchNormalization()(encoded1)

###Intermediate layer###
EbNo=10.0**(15/10.0)
channel_in = Lambda(real_to_complex)(encoded2)
channel = (channel_in)*(h) + (noise)
#channel = tf.multiply(channel_in, h) + (noise)
#channel1 = tf.multiply(channel, noise)
channel_out = Lambda(complex_to_real)(channel)

###Decoder###
decoded = Dense(M, activation='linear')(channel_out)
decoded1 = Dense(M, activation='relu')(decoded)
decoded2 = Dense(M, activation='softmax')(decoded1)

autoencoder = Model(input_data, decoded2)
#rmsprop = RMSprop(learning_rate=0.001)
sgd = SGD(learning_rate=0.02)
autoencoder.compile(optimizer='sgd', loss='categorical_crossentropy', metrics=['accuracy'])


print (autoencoder.summary())


# traning auto encoder
autoencoder.fit(data, data, epochs=100, batch_size=399)

故障原因与修正方案

核心错误点

  1. 实复转换逻辑维度错误:编码器输出encoded2维度为[batch_size, 2*n_channel](即[batch,14]),原real_to_complex函数错误取每个样本的前2个值作为实部、虚部,剩余12个维度完全丢弃,且reshape逻辑硬编码列数为7。该错误直接导致399个batch样本仅保留2个特征共798个数值,reshape为7列后行数变为114,和输入batch_size完全不匹配,就是报错里logits维度为114的直接原因。
  2. 瑞利衰落、噪声生成逻辑错误:原代码在模型图构建阶段就生成固定的h和noise张量,没有匹配batch维度,且衰落系数、噪声没有按每个样本每个符号独立生成,不符合瑞利信道的物理定义。
  3. 复转实维度错误:原complex_to_real函数拼接实部虚部后错误reshape到7维度,实际输出应为[batch_size, 2*n_channel],导致后续全连接层输入维度异常,最终输出logits维度和batch_size不匹配。
  4. 缺少依赖导入:代码中大量使用tf开头的接口,但未导入tensorflow库,运行时会直接报名称错误。

修正要点

  • 实复转换:将长度为2*n_channel的实向量,前n_channel个值作为实部,后n_channel个值作为虚部,拼接为[batch_size, n_channel]的复值张量,全程保留batch维度,不做破坏维度的硬编码reshape。
  • 信道实现:将瑞利衰落、噪声生成逻辑封装到Lambda层内部,每次前向传播时根据当前输入的batch形状动态生成对应维度的随机系数,保证每个batch、每个符号的衰落和噪声独立采样。
  • 复转实:将复值张量的实部、虚部在最后一个维度拼接,输出[batch_size, 2*n_channel]的实值张量送入解码器。
  • 补充功率归一化:编码器输出后做平均功率归一化,保证发送信号功率符合通信系统约定,避免功率不匹配导致的训练不收敛问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 17:42:43