C#实现MNIST神经网络精度升至20%后回落至10%求助
MNIST神经网络训练精度停滞问题求助
我参考Kaggle上的Python手写MNIST神经网络代码,使用C# Windows Forms实现了同款神经网络,但模型训练时精度从约9%升至15-20%后便回落至10%,无法取得有效进展,附上实现代码如下,恳请技术帮助:
namespace neuralNetwork3 { public partial class Form1 : Form { Random rnd = new Random(); string[] trainingData = File.ReadAllLines(@"file path"); string[] testingData = File.ReadAllLines(@"file path"); double[,] X_train; int m; double[,] W1 = new double[10, 784]; double[] b1 = new double[10]; double[,] W2 = new double[10, 10]; double[] b2 = new double[10]; double[,] Z1; double[,] Z2; double[,] A1; double[,] A2; public Form1() { InitializeComponent(); } private void Form1_Load(object sender, EventArgs e) { X_train = new double[trainingData.Length,784]; for (int i = 0; i < X_train.GetLength(0); i++) { double[] lineVals = trainingData[i][2..].Split(',').Select(x => double.Parse(x) / 255f).ToArray(); for (int j = 0; j < X_train.GetLength(1); j++) { X_train[i,j] = lineVals[j]; } } m = trainingData.Length; Z1 = new double[10, m]; Z2 = new double[10, m]; A1 = new double[10, m]; A2 = new double[10, m]; } void init_params() { for (int i = 0; i < W1.GetLength(0); i++) { for (int j = 0; j < W1.GetLength(1); j++) { W1[i, j] = rnd.NextDouble() - 0.5d; } } for (int i = 0; i < b1.GetLength(0); i++) { b1[i] = rnd.NextDouble() - 0.5d; } for (int i = 0; i < W2.GetLength(0); i++) { for (int j = 0; j < W2.GetLength(1); j++) { W2[i, j] = rnd.NextDouble() - 0.5d; } } for (int i = 0; i < b1.GetLength(0); i++) { b2[i] = rnd.NextDouble() - 0.5d; } } double ReLU(double Z) { return Math.Max(0, Z); } double ReLU_deriv(double Z) { return Z > 0 ? 1:0; } int[,] one_hot(int[] Y) { int[,] encoded = new int[10, trainingData.Length]; for (int i = 0; i < encoded.GetLength(0); i++) { for (int j = 0; j < encoded.GetLength(1); j++) { encoded[i, j] = Y[j] == i ? 1:0; } } return encoded; } void backward_prop(int[] Y, double learningRate) { int[,] oneHotY = one_hot(Y); double[,] dZ2 = new double[10, m]; for (int i = 0; i < dZ2.GetLength(0); i++) { for (int j = 0; j < dZ2.GetLength(1); j++) { dZ2[i, j] = 2*(A2[i,j] - oneHotY[i,j]); } } double[,] dW2 = new double[10, 10]; double[,] dZ2dotA1 = MProduct(dZ2, MTranspose(A1)); for (int i = 0; i < dW2.GetLength(0); i++) { for (int j = 0; j < dW2.GetLength(1); j++) { dW2[i, j] = (double)1 / m * dZ2dotA1[i,j]; } } double[] dB2 = new double[10]; for (int i = 0; i < dB2.Length; i++) { dB2[i] = (double)1 / m * Enumerable.Range(0, m).Select(x => dZ2[i,x]).Sum();//dZ2[i, j]; } double[,] dZ1 = new double[10, m]; double[,] W2dotdZ2 = MProduct(MTranspose(W2), dZ2); double[,] reluZ1 = new double[Z1.GetLength(0), Z1.GetLength(1)]; for (int i = 0; i < reluZ1.GetLength(0); i++) { for (int j = 0; j < reluZ1.GetLength(1); j++) { reluZ1[i, j] = ReLU_deriv(Z1[i, j]); } } dZ1 = MProduct(MTranspose(W2), dZ2); for (int i = 0; i < dZ1.GetLength(0); i++) { for (int j = 0; j < dZ1.GetLength(1); j++) { dZ1[i, j] *= reluZ1[i, j]; } } double[,] dW1 = new double[10, 784]; double[,] dZ1dotX = MProduct(dZ1, X_train); for (int i = 0; i < dW1.GetLength(0); i++) { for (int j = 0; j < dW1.GetLength(1); j++) { dW1[i, j] = (double)1 / m * dZ1dotX[i, j]; } } double[] dB1 = new double[10]; for (int i = 0; i < dB1.Length; i++) { dB1[i] = (double)1 / m * Enumerable.Range(0, m).Select(x => dZ1[i, x]).Sum();//dZ1[i, j]; } //update params for (int i = 0; i < W1.GetLength(0); i++) { for (int j = 0; j < W1.GetLength(1); j++) { W1[i, j] = W1[i,j] - learningRate * dW1[i,j]; } } for (int i = 0; i < W2.GetLength(0); i++) { for (int j = 0; j < W2.GetLength(1); j++) { W2[i, j] = W2[i, j] - learningRate * dW2[i, j]; } } for (int i = 0; i < b1.Length; i++) { b1[i] = b1[i] - learningRate * dB1[i]; } for (int i = 0; i < b2.Length; i++) { b2[i] = b2[i] - learningRate * dB2[i]; } } void gradientDescent(int[] Y, int iterations) { init_params(); for (int i = 0; i < iterations; i++) { updateIteration(i.ToString()); forward_prop(); backward_prop(Y,0.01f); testAll(i); Application.DoEvents(); } } double[] softmax(double[] Z) { double[] A = new double[Z.Length]; double[] exp = new double[10]; for (int i = 0; i < Z.Length; i++) { exp[i] = Math.Exp(Z[i] - Z.Max()); } for (int i = 0; i < Z.Length; i++) { A[i] = exp[i] / exp.Sum(); } return A; } void forward_prop() { double[,] W1dotX = MProduct(W1, MTranspose(X_train)); for (int i = 0; i < Z1.GetLength(0); i++) { for (int j = 0; j < Z1.GetLength(1); j++) { Z1[i, j] = W1dotX[i, j] + b1[i]; } } for (int i = 0; i < A1.GetLength(0); i++) { for (int j = 0; j < A1.GetLength(1); j++) { A1[i, j] = ReLU(Z1[i, j]); } } double[,] W2dotA1 = MProduct(W2, A1); for (int i = 0; i < Z2.GetLength(0); i++) { for (int j = 0; j < Z2.GetLength(1); j++) { Z2[i, j] = W2dotA1[i, j] + b2[i]; } } for (int i = 0; i < Z2.GetLength(0); i++) { double[] softZ2 = softmax(Enumerable.Range(0, Z2.GetLength(0)).Select(x => Z2[x, i]).ToArray()); for (int j = 0; j < Z2.GetLength(0); j++) { A2[i, j] = softZ2[j]; } } } double[,] MProduct(double[,] M1, double[,] M2) { double[,] M3 = new double[M1.GetLength(0), M2.GetLength(1)]; for (int i = 0; i < M1.GetLength(0); i++) { for (int j = 0; j < M2.GetLength(1); j++) { double sum = 0d; for (int k = 0; k < M1.GetLength(1); k++) { sum += M1[i, k] * M2[k, j]; } M3[i, j] = sum; } } return M3; } double[,] MTranspose(double[,] M1) { double[,] M2 = new double[M1.GetLength(1), M1.GetLength(0)]; for (int i = 0; i < M1.GetLength(0); i++) { for (int j = 0; j < M1.GetLength(1); j++) { M2[j, i] = M1[i, j]; } } return M2; } private void trainButton_Click(object sender, EventArgs e) { int[] Y = new int[trainingData.Length]; for (int i = 0; i < Y.Length; i++) { Y[i] = int.Parse(trainingData[i][0].ToString()); } gradientDescent(Y,250); MessageBox.Show("done"); } private void testButton_Click(object sender, EventArgs e) { int num; double[] O = new double[10]; if (testTextBox.Text == "*") { testAll(1); } else { try { num = int.Parse(testTextBox.Text); } catch { return; } showImg(testingData[num]); O = testExample(testingData[num][2..].Split(',').Select(x => double.Parse(x) / 255f).ToArray()); MessageBox.Show(String.Join(',', O) + " : " + Array.IndexOf(O, O.Max())); } } private void testAll(int iteration) { int num; double[] O = new double[10]; int numCorrect = 0; for (int i = 0; i < testingData.Length; i++) { O = testExample(testingData[i][2..].Split(',').Select(x => double.Parse(x) / 255f).ToArray()); if (Array.IndexOf(O, O.Max()) == int.Parse(testingData[i][0].ToString())) { numCorrect++; } } iterationLabel.Text = iteration.ToString() + " : " + ((double)numCorrect / testingData.Length * 100).ToString() + "% correct"; } double[] testExample(double[] example) { double[] H = new double[10]; double [] O = new double[10]; for (int i = 0; i < W1.GetLength(0); i++) { double sum = 0f; for (int j = 0; j < W1.GetLength(1); j++) { sum += W1[i,j] * example[j]; } H[i] = sum + b1[i]; } for (int i = 0; i < W2.GetLength(0); i++) { double sum = 0f; for (int j = 0; j < W2.GetLength(1); j++) { sum += W2[i, j] * H[j]; } O[i] = sum + b2[i]; } return O; } async void showImg(string line) { //MessageBox.Show(trainingData[0]); Bitmap bmp = new Bitmap(28, 28); List<int> pixelBrightnesses = new List<int>(); pixelBrightnesses = line.Split(',').Select(int.Parse).ToList(); for (int i = 0; i < 28; i++) { for (int j = 0; j < 28; j++) { int colour = int.Parse(pixelBrightnesses[j * 28 + i + 1].ToString()); bmp.SetPixel(i, j, Color.FromArgb(colour, colour, colour)); } } pictureBox1.Image = bmp; await Task.Delay(500); } async void updateIteration(string iteration) { //MessageBox.Show(iteration.ToString()); iterationLabel.Text = "iteration: " + iteration; await Task.Delay(500); } } }
内容的提问来源于stack exchange,提问作者Louis
相关产品推荐
相关产品推荐

