基于MXNet二次方程求解模型出现ZeroDivisionError问题求助
Let's break down your problem into two parts: fixing the immediate ZeroDivisionError, and then improving your model to actually solve quadratic equations effectively.
1. Fixing the ZeroDivisionError
The error comes directly from the mx.callback.Speedometer you're using. Here's why it happens:
speed = self.frequent * self.batch_size / (time.time() - self.tic) ZeroDivisionError: float division by zero
When your batch_size=1, each batch processes almost instantly—so the time difference between the callback's start (self.tic) and current time is so small it rounds to 0. The Speedometer tries to divide by this zero value, causing the crash.
Here are 3 quick fixes:
- Increase your batch size: Use a larger batch size (like 10 or 20) so each batch takes measurable time to process.
- Adjust the Speedometer's frequency: Change the second parameter in
Speedometerto a higher number (e.g.,mx.callback.Speedometer(batch_size, 5)), so it calculates speed less often. - Remove the Speedometer temporarily: If you don't need the speed metrics right now, just delete the
batch_end_callbackargument frommodel.fit().
For example, modifying your fit call with a larger batch size:
batch_size = 10 train_iter = mx.io.NDArrayIter(train_data,train_label, batch_size, shuffle=True,label_name='lin_reg_label') eval_iter = mx.io.NDArrayIter(eval_data, eval_label, batch_size, shuffle=False) # ... rest of your code ... model.fit(train_iter, eval_iter, optimizer_params={'learning_rate':0.005, 'momentum': 0.9}, num_epoch=50, eval_metric='mse', batch_end_callback = mx.callback.Speedometer(batch_size, 2))
2. Improving Your Quadratic Equation Solver Model
Even after fixing the error, your current linear regression model won't work well for this task. The root of a quadratic equation (-b + sqrt(b²-4ac))/(2a) is a non-linear function of a, b, c, but a simple fully connected layer with linear output can't capture non-linear relationships. That's why you're seeing exploding MSE values during training.
Here are two better approaches:
Option 1: Use a Non-Linear Neural Network
Add hidden layers with activation functions to model the non-linear relationship. For example:
X = mx.sym.Variable('data') Y = mx.symbol.Variable('lin_reg_label') # Add hidden layers with ReLU activation to capture non-linearity fc1 = mx.sym.FullyConnected(data=X, name='fc1', num_hidden=16) act1 = mx.sym.Activation(data=fc1, name='act1', act_type="relu") fc2 = mx.sym.FullyConnected(data=act1, name='fc2', num_hidden=8) act2 = mx.sym.Activation(data=fc2, name='act2', act_type="relu") fc3 = mx.sym.FullyConnected(data=act2, name='fc3', num_hidden=1) lro = mx.sym.LinearRegressionOutput(data=fc3, label=Y, name="lro") model = mx.mod.Module( symbol = lro , data_names=['data'], label_names = ['lin_reg_label'])
Option 2: Engineer Features to Fit Linear Model
Since the root formula includes terms like b² and a*c, you can precompute these features and feed them into a linear model. This lets the linear model learn the coefficients of the root formula directly:
# Preprocess training data to include non-linear features matching the root formula def preprocess_features(data): a = data[:,0] b = data[:,1] c = data[:,2] # Add features: b², a*c, b/a, c/a (since root involves division by 2a) features = np.column_stack([a, b, c, b**2, a*c, b/a, c/a]) return features train_data_processed = preprocess_features(train_data) eval_data_processed = preprocess_features(eval_data) # Use processed data in iterators train_iter = mx.io.NDArrayIter(train_data_processed,train_label, batch_size, shuffle=True,label_name='lin_reg_label') eval_iter = mx.io.NDArrayIter(eval_data_processed, eval_label, batch_size, shuffle=False) # Keep your original linear model—it will now fit the engineered features X = mx.sym.Variable('data') Y = mx.symbol.Variable('lin_reg_label') fully_connected_layer = mx.sym.FullyConnected(data=X, name='fc1', num_hidden = 1) lro = mx.sym.LinearRegressionOutput(data=fully_connected_layer, label=Y, name="lro")
Final Notes
- For your training data, multiplying
bby 25 to avoid non-real roots is a smart move—keep that logic in place. - If your MSE is still exploding, try lowering the learning rate (e.g., 0.001 instead of 0.005) to stabilize training.
内容的提问来源于stack exchange,提问作者Eng1234

