Keras实现整数乘法神经网络不收敛的原因及适配方案咨询
Hey there! I’ve run into similar issues when training neural networks on multiplicative tasks before, so let’s break down what’s going on here and how to fix it.
1. Why the multiplication model fails to converge
Let’s start with the root causes:
Massive scale mismatch between input and output: For your task, inputs are integers between 0-500. Addition outputs max out at 1000, but multiplication outputs can go up to 250,000. The Mean Squared Error (MSE) loss gets dominated by these large values, making it impossible for the Adam optimizer to adjust weights effectively—think of trying to fine-tune a tiny knob when the error signal is screaming at 7 billion.
Task complexity difference: Addition is a linear operation, which even a simple feedforward network can learn easily. Multiplication is a bilinear operation (it depends on the product of two inputs), and vanilla dense layers struggle to capture this interaction without explicit hints or proper data scaling.
Suboptimal activation function choice: Using
reluon the final layer isn’t a dealbreaker here, but combined with the scale issue, it can slow down convergence. Since your outputs are non-negative, relu works, but a linear activation (activation='linear') is more natural for regression tasks with unbounded outputs.
2. Fixes to train a working multiplication model
Here are concrete changes you can make to your code:
Step 1: Normalize your data
Scaling both inputs and outputs to a small range (like 0-1) is the single most impactful fix. Use scikit-learn’s MinMaxScaler to handle this:
from sklearn.preprocessing import MinMaxScaler import numpy as np from keras.models import Sequential from keras.layers import Dense def create_data(low, high, examples): train_data = [] label_data = [] a = np.random.randint(low=low, high=high, size=examples, dtype='int') b = np.random.randint(low=low, high=high, size=examples, dtype='int') for i in range(examples): train_data.append([a[i], b[i]]) label_data.append(a[i] * b[i]) return np.array(train_data), np.array(label_data) # Generate data X, y = create_data(0, 500, 10000) y = y.reshape(-1, 1) # Reshape for scaler compatibility # Scale inputs and outputs to 0-1 range scaler_X = MinMaxScaler(feature_range=(0, 1)) scaler_y = MinMaxScaler(feature_range=(0, 1)) X_scaled = scaler_X.fit_transform(X) y_scaled = scaler_y.fit_transform(y)
Step 2: Adjust the network structure and hyperparameters
- Swap the final
reluactivation forlinear(since we’re doing regression, not classification). - Optionally, add an explicit product feature to help the network learn the bilinear relationship directly:
model = Sequential() # Input now includes [a, b, a*b] to give the network a hint about the multiplicative relationship model.add(Dense(8, input_dim=3)) model.add(Dense(16, activation='relu')) model.add(Dense(8, activation='relu')) model.add(Dense(1, activation='linear')) # Linear output for unbounded regression # Compile with default Adam, or tweak learning rate if needed (e.g., optimizer=Adam(learning_rate=0.0005)) model.compile(optimizer='adam', loss='mean_squared_error') # Prepare augmented input with the product feature X_augmented = np.hstack([X_scaled, (X_scaled[:,0] * X_scaled[:,1]).reshape(-1,1)]) # Train with more epochs for the complex task model.fit(X_augmented, y_scaled, epochs=50, batch_size=32, verbose=1)
Step 3: Inference with inverse scaling
When making predictions, don’t forget to scale the results back to the original integer range:
test_data, _ = create_data(0, 500, 10) test_scaled = scaler_X.transform(test_data) test_augmented = np.hstack([test_scaled, (test_scaled[:,0] * test_scaled[:,1]).reshape(-1,1)]) results_scaled = model.predict(test_augmented) results = scaler_y.inverse_transform(results_scaled) # Calculate and print errors for i in range(len(test_data)): true_val = test_data[i][0] * test_data[i][1] print(f'test value: {test_data[i]}, true product: {true_val}, predicted: {results[i][0]:.2f}, error: {results[i][0] - true_val:.2f}')
Optional: Try alternative loss functions
If MSE still gives trouble, switch to Mean Absolute Error (loss='mean_absolute_error'), which is less sensitive to extreme outliers in the output values.
Testing the fix
With these changes, you’ll see the loss drop steadily each epoch, and the model will start making accurate predictions for integer multiplication tasks.
内容的提问来源于stack exchange,提问作者D.Griffiths

