用JAVA进行神经网络建模及泛化能力测试

from http://slowman.iteye.com/blog/722602

题目:

                     表一: 澳大利亚野兔眼睛晶状体重量与年龄的对应关系

 

 

编号

年龄(天)

重量(mg)

年龄(天)

重量(mg)

年龄(天)

重量(mg)

年龄(天)

重量(mg)

1

2

3

4

5

6

7

8

9

10

11

12

13

14

15

16

17

18

15

15

15

18

28

29

37

37

44

50

50

60

61

64

65

65

72

75

21.66

22.75

22.3

31.25

44.79

40.55

50.25

46.88

52.03

63.47

61.13

81

73.09

79.09

79.51

65.31

71.9

86.1

75

82

85

91

91

97

98

125

142

142

147

147

150

159

165

183

192

195

94.6

92.5

105

101.7

102.9

110

104.3

134.9

130.68

140.58

155.3

152.2

144.5

142.15

139.81

153.22

145.72

161.1

218

218

219

224

225

227

232

232

237

246

258

276

285

300

301

305

312

317

174.18

173.03

173.54

178.86

177.68

173.73

159.98

161.29

187.07

176.13

183.4

186.26

189.66

186.09

186.7

186.8

195.1

216.41

338

347

354

357

375

394

513

535

554

591

648

660

705

723

756

768

860

203.23

188.38

189.7

195.31

202.63

224.82

203.3

209.7

233.9

234.7

244.3

231

242.4

230.77

242.57

232.12

246.7

 

澳大利亚野兔眼睛晶状体的重量为年龄的函数。利用BP算法,设计一个多层感知器,为表中的数据集提供一个非线性逼近,并测试其泛化能力。

 

算法源码:

 

Java代码  收藏代码
  1. package com.lwm.cn.althom;  
  2.   
  3. import java.io.BufferedReader;  
  4. import java.io.BufferedWriter;  
  5. import java.io.File;  
  6. import java.io.FileNotFoundException;  
  7. import java.io.FileReader;  
  8. import java.io.FileWriter;  
  9. import java.io.IOException;  
  10. import java.text.SimpleDateFormat;  
  11. import java.util.ArrayList;  
  12. import java.util.Date;  
  13. import java.util.GregorianCalendar;  
  14. import java.util.Random;  
  15.   
  16. public class BackProp {  
  17.     private int randomPrecision = 8; // 生成double型随机数的精度,默认为6位小数  
  18.   
  19.     private int input_dimension; // 输入向量的维度  
  20.   
  21.     private int output_dimension; // 输出向量的维数  
  22.   
  23.     private int mid_dimension; // 隐层结点的个数  
  24.   
  25.     private double[][] V; // 输入层到隐层的权值矩阵  
  26.   
  27.     private double[][] W; // 隐层到输出层的权值矩阵  
  28.   
  29.     private double[] inputArray; // 输入层向量  
  30.   
  31.     private double[] midArray; // 隐层输出向量  
  32.   
  33.     private double[] outputArray; // 输出层向量  
  34.   
  35.     private double[] teacherArray; // 期望层向量  
  36.   
  37.     private double mid_Threshold; // 隐层阈值  
  38.   
  39.     private double out_Threshold; // 输出层阈值  
  40.   
  41.     private double[] midError; // 隐层的误差  
  42.   
  43.     private double[] outError; // 输出层的误差  
  44.   
  45.     private double totalError = 0.0;  
  46.   
  47.     private double outPrecision; // 要达到的精度  
  48.   
  49.     private double learnRate; // 学习的速率  
  50.   
  51.     private int trainTotal = 3000; // 学习1000次  
  52.   
  53.     private boolean isQualify = false; // 用于判断是不是达到精度要求  
  54.   
  55.     private ArrayList<SampleNode> trainArray = new ArrayList<SampleNode>(100); // 存入训练集  
  56.   
  57.     private ArrayList<SampleNode> testArray = new ArrayList<SampleNode>(100); // 存放测试集  
  58.   
  59.     private BufferedWriter bw = null; // 用于将学习和测试过程写于文件  
  60.   
  61.     Date startTime;  
  62.   
  63.     // SampleNode sample;  
  64.   
  65.     // Math.random()  
  66.     public BackProp(double[][] v, double[][] w, int input_dimension,  
  67.             int output_dimension, int mid_dimension) {  
  68.         super();  
  69.         V = v;  
  70.         W = w;  
  71.         this.input_dimension = input_dimension;  
  72.         this.output_dimension = output_dimension;  
  73.         this.mid_dimension = mid_dimension;  
  74.     }  
  75.   
  76.     /** 
  77.      * 默认构造函数 ,对于本次实验,输入向量只有一个,输出也只有一个. 隐层结点的个数默认为4 
  78.      *  
  79.      */  
  80.     public BackProp() {  
  81.         input_dimension = 1;  
  82.         output_dimension = 1;  
  83.         mid_dimension = 8;  
  84.   
  85.         inputArray = new double[input_dimension];  
  86.         teacherArray = new double[output_dimension];  
  87.         midArray = new double[mid_dimension];  
  88.         outputArray = new double[output_dimension];  
  89.   
  90.         V = new double[input_dimension][mid_dimension];  
  91.         W = new double[mid_dimension][input_dimension];  
  92.   
  93.         midError = new double[mid_dimension];  
  94.         outError = new double[output_dimension];  
  95.     }  
  96.   
  97.     /** 
  98.      * 初始化函数 我认为一个完整的BP算法应该具备通用性,可以任意设置输入结点个数和隐层的层数及每一层的结点个数 
  99.      * 初始化权值矩阵V和W,每个元素的值均为0-1之间的六位小数 
  100.      */  
  101.   
  102.     public void init()  
  103.     {  
  104.         // 记录程序开始时间及结束时间,以开始时间命名一个文件,用来保存学习和测试结果.  
  105.         startTime = new Date();  
  106.         SimpleDateFormat sdf = new SimpleDateFormat("yyyy年MM月dd日HH时mm分ss秒");  
  107.         String timeStr = sdf.format(startTime);  
  108.         String filePathName = "E:" + File.separator + timeStr + ".txt";  
  109.         try  
  110.         {  
  111.             bw = new BufferedWriter(new FileWriter(filePathName));  
  112.             bw.write("程序开始时间:" + timeStr + "\n");  
  113.         } catch (IOException e)  
  114.         {  
  115.             // TODO Auto-generated catch block  
  116.             e.printStackTrace();  
  117.         }  
  118.         mid_Threshold = MathExtend.round(Math.random(), randomPrecision); // 初始化隐层的阈值  
  119.         out_Threshold = MathExtend.round(Math.random(), randomPrecision); // 初始化输出层的阈值  
  120.         // 初始化V矩阵  
  121.         for (int i = 0; i < input_dimension; i++)  
  122.             for (int j = 0; j < mid_dimension; j++)  
  123.                 V[i][j] = MathExtend.round(Math.random(), randomPrecision);  
  124.   
  125.         // 初始化W矩阵  
  126.         for (int i = 0; i < mid_dimension; i++)  
  127.             for (int j = 0; j < output_dimension; j++)  
  128.                 W[i][j] = MathExtend.round(Math.random(), randomPrecision);  
  129.   
  130.         // 置总的误差为0,学习率为0-1之间的小数,网络训练后达到的精度为一正小数  
  131.         totalError = 0.0;  
  132.         learnRate = MathExtend.round(Math.random(), randomPrecision);  
  133.         // learnRate = 0.12;  
  134.         outPrecision = MathExtend.round(Math.random(), randomPrecision);  
  135.   
  136.         try  
  137.         {  
  138.             StringBuilder sb = new StringBuilder();  
  139.             sb.append("本次实验随机生成的学习率: " + learnRate);  
  140.             sb.append("\n");  
  141.             sb.append("期望达到的精度为: " + outPrecision);  
  142.             sb.append("\n");  
  143.             bw.write(sb.toString());  
  144.         } catch (IOException e)  
  145.         {  
  146.             // TODO Auto-generated catch block  
  147.             e.printStackTrace();  
  148.         }  
  149.   
  150.         getTrainData(); // 取得训练集  
  151.         getTestData(); // 取得测试集  
  152.         normalized(); // 归一化  
  153.     }  
  154.   
  155.     /** 
  156.      * @author Administrator 输入层向隐层,隐层向输出层的传播 
  157.      *  
  158.      */  
  159.     public void finish()  
  160.     {  
  161.         // Date endDate = new Date() ;  
  162.   
  163.         try  
  164.         {  
  165.             bw.close();  
  166.         } catch (IOException e)  
  167.         {  
  168.             // TODO Auto-generated catch block  
  169.             e.printStackTrace();  
  170.         }  
  171.     }  
  172.   
  173.     public void forword()  
  174.     {  
  175.         int i, j;  
  176.         double temp_sum ; // 用于向量的内积  
  177.         // 输出层到隐层  
  178.         for (i = 0; i < mid_dimension; i++)  
  179.         {  
  180.             temp_sum = 0.0  ; //初始化为0  
  181.             for (j = 0; j < input_dimension; j++)  
  182.                 temp_sum += V[j][i] * inputArray[j];  
  183.             temp_sum = temp_sum - mid_Threshold;  
  184.             midArray[i] = 1.0 / (1 + Math.exp(-temp_sum));  
  185.         }  
  186.   
  187.           
  188.         // 隐层到输出层  
  189.         for (i = 0; i < output_dimension; i++)  
  190.         {  
  191.             temp_sum = 0.0; // 初始化  
  192.             for (j = 0; j < mid_dimension; j++)  
  193.                 temp_sum = W[j][i] * midArray[j];  
  194.             temp_sum = temp_sum - out_Threshold;  
  195.             outputArray[i] = 1.0 / (1 + Math.exp(-temp_sum));  
  196.         }  
  197.         // 计算误差,累加起来,  
  198.         temp_sum = 0.0;  
  199.         for (i = 0; i < output_dimension; i++)  
  200.         {  
  201.             temp_sum = teacherArray[i] - outputArray[i]; // 注意中,本设计中output_dimension=1的  
  202.             totalError += temp_sum * temp_sum / 2;  
  203.         }  
  204.         // printResult();  
  205.     }  
  206.   
  207.     private void printResult()  
  208.     {  
  209.         /* 
  210.          * StringBuilder sb = new StringBuilder() ; 
  211.          * sb.append("输入数据:"+inputArray[0]); sb.append(" 
  212.          * 实际输出数据:"+outputArray[0]); sb.append(" 期望输出数据为:"+teacherArray[0]) ; 
  213.          * sb.append("\\n") ; try { bw.write(sb.toString()); } catch 
  214.          * (IOException e) { // TODO Auto-generated catch block 
  215.          * e.printStackTrace(); } 
  216.          */  
  217.         System.out.print("输入数据:" + inputArray[0]);  
  218.         System.out.print("   实际输出数据:" + outputArray[0]);  
  219.         System.out.println("   期望输出数据为:" + teacherArray[0]);  
  220.     }  
  221.   
  222.     /** 
  223.      * 反向调整权值矩阵 
  224.      */  
  225.     public void adjustWeight()  
  226.     {  
  227.         double temp_sum = 0.0;  
  228.         int i, j;  
  229.         // 计算各层的误差信号  输出层  
  230.         for (i = 0; i < output_dimension; i++)  
  231.         {  
  232.             outError[i] = (teacherArray[i] - outputArray[i])  
  233.                     * (1 - outputArray[i]) * outputArray[i];  
  234.         }  
  235. //    隐层误差  
  236.         for (i = 0; i < mid_dimension; i++)  
  237.         {  
  238.             temp_sum=0.0d ;  
  239.             for (j = 0; j < output_dimension; j++)  
  240.                 temp_sum += outError[j] * W[i][j];  
  241.             midError[i] = temp_sum * (1 - midArray[i]) * midArray[i];  
  242.         }  
  243.   
  244.         // 调整W权值矩阵  
  245.         for (i = 0; i < mid_dimension; i++)  
  246.         {  
  247.             for (j = 0; j < output_dimension; j++)  
  248.                 W[i][j] += learnRate * outError[j] * midArray[i];  
  249.         }  
  250.         // 调整V权值矩阵  
  251.   
  252.         for (i = 0; i < input_dimension; i++)  
  253.             for (j = 0; j < mid_dimension; j++)  
  254.                 V[i][j] += learnRate * midError[j] * inputArray[i];  
  255.   
  256.     }  
  257.   
  258.     public void getTrainData()  
  259.     {  
  260.         String filePathName = "E:" + File.separator + "traindata.txt";  
  261.         BufferedReader br = null;  
  262.         try  
  263.         {  
  264.             br = new BufferedReader(new FileReader(filePathName));  
  265.         } catch (FileNotFoundException e)  
  266.         {  
  267.             // TODO Auto-generated catch block  
  268.             e.printStackTrace();  
  269.         }  
  270.         String s = null;  
  271.         SampleNode sNode = null;  
  272.         try  
  273.         {  
  274.             while ((s = br.readLine()) != null)  
  275.             {  
  276.                 String data[] = s.trim().split("[\\s]+");  
  277.                 if (data == null || data.length != 2)  
  278.                 {  
  279.   
  280.                     System.out.println("traindata文件数据有问题!");  
  281.                     return;  
  282.                 }  
  283.                 double in = Double.parseDouble(data[0]);  
  284.                 double hope = Double.parseDouble(data[1]);  
  285.                 sNode = new SampleNode(in, hope);  
  286.                 trainArray.add(sNode);  
  287.   
  288.                 // trainArray.  
  289.             }  
  290.         } catch (IOException e)  
  291.         {  
  292.             // TODO Auto-generated catch block  
  293.             e.printStackTrace();  
  294.         }  
  295.         trainArray.trimToSize();  
  296.     }  
  297.   
  298.     public void getTestData()  
  299.     {  
  300.         String fileName = "E:" + File.separator + "testdata.txt";  
  301.         BufferedReader br = null;  
  302.         try  
  303.         {  
  304.             br = new BufferedReader(new FileReader(fileName));  
  305.         } catch (FileNotFoundException e)  
  306.         {  
  307.             // TODO Auto-generated catch block  
  308.             System.out.println("testdata.txt文件不存在");  
  309.             e.printStackTrace();  
  310.         }  
  311.   
  312.         String s = null;  
  313.         SampleNode sNode = null;  
  314.         try  
  315.         {  
  316.             while ((s = br.readLine()) != null)  
  317.             {  
  318.                 String data[] = s.trim().split("[\\s]+");  
  319.                 if (data == null || data.length != 2)  
  320.                 {  
  321.   
  322.                     System.out.println("testdata文件数据有问题!");  
  323.                     return;  
  324.                 }  
  325.                 double in = Double.parseDouble(data[0]);  
  326.                 double hope = Double.parseDouble(data[1]);  
  327.                 sNode = new SampleNode(in, hope);  
  328.                 testArray.add(sNode);  
  329.   
  330.                 // trainArray.  
  331.             }  
  332.         } catch (IOException e)  
  333.         {  
  334.             // TODO Auto-generated catch block  
  335.             e.printStackTrace();  
  336.         }  
  337.         testArray.trimToSize();  
  338.   
  339.     }  
  340.   
  341.     /** 
  342.      * 对输入数据进行归一化处理,将输入数据限制在[0,1]区间内 
  343.      *  
  344.      */  
  345.   
  346.     private void normalized()  
  347.     {  
  348.         if (trainArray == null || trainArray.size() == 0 || testArray == null  
  349.                 || testArray.size() == 0)  
  350.         {  
  351.             System.out.println("测试数据或者训练数据有问题!");  
  352.             return;  
  353.         }  
  354.         SampleNode sNode = null;  
  355.         // 训练数据归一化  
  356.         int size = trainArray.size();  
  357.         int i = 0;  
  358.         while (i < size)  
  359.         {  
  360.             sNode = trainArray.get(i);  
  361.             double in = sNode.in;  
  362.             double hope = sNode.hope;  
  363.             in /= 1000.0; // 归一  
  364.             hope /= 250.0;  
  365.             sNode.in = in;  
  366.             sNode.hope = hope;  
  367.             trainArray.set(i, sNode);  
  368.             i++;  
  369.         }  
  370.   
  371.         size = testArray.size();  
  372.         i = 0;  
  373.         // 测试数据归一化  
  374.         while (i < size)  
  375.         {  
  376.             sNode = testArray.get(i);  
  377.             double in = sNode.in;  
  378.             double hope = sNode.hope;  
  379.             in /= 1000.0; // 归一  
  380.             hope /= 250.0;  
  381.             sNode.in = in;  
  382.             sNode.hope = hope;  
  383.             trainArray.set(i, sNode);  
  384.             i++;  
  385.         }  
  386.   
  387.     }  
  388.   
  389.     public void startTrain()  
  390.     {  
  391.         if (trainArray == null || trainArray.size() == 0)  
  392.             return;  
  393.         System.out.println("训练开始");  
  394.         System.out.println("当前学习速率:" + learnRate);  
  395.         System.out.println("期望精度为:" + outPrecision);  
  396.         int trainConunter = 0;  
  397.         while (trainConunter++ < trainTotal)  
  398.         {  
  399.             System.out.println("第" + trainConunter + "次训练开始:");  
  400.             for (SampleNode sNode : trainArray)  
  401.             {  
  402.                 // 说明:在本设计中inputArray,和teacherArray虽然都是数组,但均只有一个元素.  
  403.                 // 本人为了综合虑,才将设为数组的.  
  404.                 inputArray[0] = sNode.in;  
  405.                 teacherArray[0] = sNode.hope;  
  406.                 forword(); // 学习一次  
  407.                 printResult();  
  408.             } // 至此,所有训练集全部学习完毕,下面应该进行权值调整.  
  409.             /*System.out.println("此次学习后,总的误差为:" + totalError); 
  410.             StringBuilder sb = new StringBuilder(); 
  411.             sb.append("第" + trainConunter); 
  412.             sb.append("次学习后,总的误差为:" + totalError); 
  413.             sb.append("\n");*/  
  414.             try  
  415.             {  
  416.             //  bw.write(sb.toString());  
  417.                 bw.write(Double.toString(totalError)+"\n") ;  
  418.             } catch (IOException e)  
  419.             {  
  420.                 e.printStackTrace();  
  421.             }  
  422.             adjustWeight(); // 集体主义原则来调整权值  
  423.   
  424.             if (totalError <= outPrecision)  
  425.             {  
  426.                 isQualify = true; // 置标志位为真,表示达到要求  
  427.   
  428.                 break;  
  429.             }  
  430.             totalError = 0.0; // 误差初化  
  431.         }  
  432.   
  433.         Date endTime = new Date();  
  434.         SimpleDateFormat sdf = new SimpleDateFormat("yyyy年MM月dd日HH时mm分ss秒");  
  435.         String endtimeStr = sdf.format(endTime);  
  436.         long gap = endTime.getTime() - this.startTime.getTime();  
  437.         StringBuilder sb = new StringBuilder();  
  438.         try  
  439.         {  
  440.             sb.append("训练结束时间为:" + endtimeStr);  
  441.             sb.append("\n");  
  442.             sb.append("总的学习时间为:" + gap);  
  443.             sb.append("微秒\n");  
  444.             sb.append("********************************************\n");  
  445.             bw.write(sb.toString());  
  446.         } catch (IOException e1)  
  447.         {  
  448.             // TODO Auto-generated catch block  
  449.             e1.printStackTrace();  
  450.         }  
  451.   
  452.         if (!isQualify)  
  453.         {  
  454.             System.out.println("达到训练次数,训练结束!");  
  455.             try  
  456.             {  
  457.                 bw.write("训练次数:" + trainTotal + "次\n");  
  458.             } catch (IOException e)  
  459.             {  
  460.                 // TODO Auto-generated catch block  
  461.                 e.printStackTrace();  
  462.             }  
  463.         } else  
  464.         {  
  465.             try  
  466.             {  
  467.                 bw.write("达到精度要求,学习完毕!\n");  
  468.             } catch (IOException e)  
  469.             {  
  470.                 // TODO Auto-generated catch block  
  471.                 e.printStackTrace();  
  472.             }  
  473.             System.out.println("达到要求的精度,训练结束!");  
  474.         }  
  475.     }  
  476.   
  477.     public void startTest()  
  478.     {  
  479.   
  480.         if (testArray == null || testArray.isEmpty() == true)  
  481.             return;  
  482.           
  483.         for (SampleNode sNode : testArray)  
  484.         {  
  485.             StringBuilder sb = new StringBuilder();  
  486.             inputArray[0] = sNode.in;  
  487.             teacherArray[0] = sNode.hope;  
  488.             forword();  
  489.             sb.append("输入测试数据: " + inputArray[0]);  
  490.             sb.append("   实际输出:" + outputArray[0]);  
  491.             sb.append("   期望输出:" + teacherArray[0]);  
  492.             sb.append("\n");  
  493.             try  
  494.             {  
  495.                 bw.write(sb.toString());  
  496.             } catch (IOException e)  
  497.             {  
  498.                 // TODO Auto-generated catch block  
  499.                 e.printStackTrace();  
  500.             }  
  501.             printResult();  
  502.         }  
  503.     }  
  504.   
  505. }  

 

 

 

 

测试输出结果如下图:

 

 

 

 

程序运行一次的收敛图如下图:


posted @ 2017-08-01 01:40  princessd8251  阅读(320)  评论(0)    收藏  举报