一层简单人工神经网络的Java实现

 From  http://blog.csdn.net/csj941227/article/details/73325695

2、数据类

 

[java] view plain copy
 
  1. import java.util.Arrays;  
  2.   
  3. public class Data {  
  4.     double[] vector;  
  5.     int dimention;  
  6.     int type;  
  7.     public double[] getVector() {  
  8.         return vector;  
  9.     }  
  10.     public void setVector(double[] vector) {  
  11.         this.vector = vector;  
  12.     }  
  13.     public int getDimention() {  
  14.         return dimention;  
  15.     }  
  16.     public void setDimention(int dimention) {  
  17.         this.dimention = dimention;  
  18.     }  
  19.     public int getType() {  
  20.         return type;  
  21.     }  
  22.     public void setType(int type) {  
  23.         this.type = type;  
  24.     }  
  25.     public Data(double[] vector, int dimention, int type) {  
  26.         super();  
  27.         this.vector = vector;  
  28.         this.dimention = dimention;  
  29.         this.type = type;  
  30.     }  
  31.     public Data() {  
  32.     }  
  33.     @Override  
  34.     public String toString() {  
  35.         return "Data [vector=" + Arrays.toString(vector) + ", dimention=" + dimention + ", type=" + type + "]";  
  36.     }  
  37.       
  38. }  

3、简单人工神经网络

 

 

[java] view plain copy
 
  1. package cn.edu.hbut.chenjie;  
  2.   
  3. import java.util.ArrayList;  
  4. import java.util.List;  
  5. import java.util.Random;  
  6.   
  7. import org.jfree.chart.ChartFactory;  
  8. import org.jfree.chart.ChartFrame;  
  9. import org.jfree.chart.JFreeChart;  
  10. import org.jfree.data.xy.DefaultXYDataset;  
  11. import org.jfree.ui.RefineryUtilities;  
  12.   
  13.   
  14. public class ANN2 {  
  15.     private double eta;//学习率  
  16.     private int n_iter;//权重向量w[]训练次数  
  17.     private List<Data> exercise;//训练数据集  
  18.     private double w0 = 0;//阈值  
  19.     private double x0 = 1;//固定值  
  20.     private double[] weights;//权重向量,其长度为训练数据维度+1,在本例中数据为2维,故长度为3  
  21.     private int testSum = 0;//测试数据总数  
  22.     private int error = 0;//错误次数  
  23.     DefaultXYDataset xydataset = new DefaultXYDataset();  
  24.       
  25.     /** 
  26.      * 向图表中增加同类型的数据 
  27.      * @param type 类型 
  28.      * @param a 所有数据的第一个分量 
  29.      * @param b 所有数据的第二个分量 
  30.      */  
  31.     public void add(String type,double[] a,double[] b)  
  32.     {  
  33.         double[][] data = new double[2][a.length];  
  34.         for(int i=0;i<a.length;i++)  
  35.         {  
  36.             data[0][i] = a[i];  
  37.             data[1][i] = b[i];  
  38.         }  
  39.         xydataset.addSeries(type, data);    
  40.     }  
  41.       
  42.     /** 
  43.      * 画图 
  44.      */  
  45.     public void draw()  
  46.     {  
  47.         JFreeChart jfreechart = ChartFactory.createScatterPlot("exercise", "x1", "x2", xydataset);  
  48.         ChartFrame frame = new ChartFrame("训练数据", jfreechart);  
  49.         frame.pack();  
  50.         RefineryUtilities.centerFrameOnScreen(frame);  
  51.         frame.setVisible(true);  
  52.     }  
  53.       
  54.     public static void main(String[] args)  
  55.     {  
  56.         ANN2 ann2 = new ANN2(0.001,100);//构造人工神经网络  
  57.           
  58.         List<Data> exercise = new ArrayList<Data>();//构造训练集  
  59.           
  60.         //人工模拟1000条训练数据 ,分界线为x2=x1+0.5  
  61.         for(int i=0;i<1000000;i++)  
  62.         {  
  63.             Random rd = new Random();  
  64.             double x1 = rd.nextDouble();//随机产生一个分量  
  65.             double x2 = rd.nextDouble();//随机产生另一个分量  
  66.             double[] da = {x1,x2};//产生数据向量  
  67.             Data d = new Data(da, 2, x2 > x1+0.5 ? 1 : -1);//构造数据  
  68.             exercise.add(d);//将训练数据加入训练集  
  69.         }  
  70.           
  71.         int sum1 = 0;//记录类型1的训练记录数  
  72.         int sum2 = 0;//记录类型-1的训练记录数  
  73.         for(int i = 0; i < exercise.size(); i++)  
  74.         {  
  75.             if(exercise.get(i).getType()==1)  
  76.                 sum1++;  
  77.             else if(exercise.get(i).getType()==-1)  
  78.                 sum2++;  
  79.         }  
  80.         double[] x1 = new double[sum1];  
  81.         double[] y1 = new double[sum1];  
  82.         double[] x2 = new double[sum2];  
  83.         double[] y2 = new double[sum2];  
  84.         int index1 = 0;  
  85.         int index2 = 0;  
  86.         for(int i = 0; i < exercise.size(); i++)  
  87.         {  
  88.             if(exercise.get(i).getType()==1)  
  89.             {  
  90.                 x1[index1] = exercise.get(i).vector[0];  
  91.                 y1[index1++] = exercise.get(i).vector[1];  
  92.             }  
  93.             else if(exercise.get(i).getType()==-1)  
  94.             {  
  95.                 x2[index2] = exercise.get(i).vector[0];  
  96.                 y2[index2++] = exercise.get(i).vector[1];  
  97.             }  
  98.         }  
  99.           
  100.         ann2.add("1", x1, y1);  
  101.         ann2.add("-1", x2, y2);  
  102.         ann2.draw();  
  103.           
  104.         ann2.input(exercise);//将训练集输入人工神经网络  
  105.           
  106.         ann2.fit();//训练  
  107.           
  108.         ann2.showWeigths();//显示权重向量  
  109.           
  110.           
  111.         //人工生成一千条测试数据  
  112.         for(int i=0;i<10000;i++)  
  113.         {  
  114.             Random rd = new Random();  
  115.             double x1_ = rd.nextDouble();  
  116.             double x2_ = rd.nextDouble();  
  117.             double[] da = {x1_,x2_};  
  118.             Data test = new Data(da, 2, x2_ > x1_+0.5 ? 1 : -1);  
  119.             ann2.predict(test);//测试  
  120.         }  
  121.           
  122.         System.out.println("总共测试" + ann2.testSum + "条数据,有" + ann2.error + "条错误,错误率:" + ann2.error * 1.0 /ann2.testSum * 100 + "%");  
  123.     }  
  124.       
  125.     /** 
  126.      *  
  127.      * @param eta 学习率 
  128.      * @param n_iter 权重分量学习次数 
  129.      */  
  130.     public ANN2(double eta, int n_iter) {  
  131.         this.eta = eta;  
  132.         this.n_iter = n_iter;  
  133.     }  
  134.   
  135.   
  136.     /** 
  137.      * 输入训练集到人工神经网络 
  138.      * @param exercise 
  139.      */  
  140.     private void input(List<Data> exercise) {  
  141.         this.exercise = exercise;//保存训练集  
  142.         weights = new double[exercise.get(0).dimention + 1];//初始化权重向量,其长度为训练数据维度+1  
  143.         weights[0] = w0;//权重向量第一个分量为w0  
  144.         for(int i = 1; i < weights.length; i++)  
  145.             weights[i] = 0;//其余分量初始化为0  
  146.     }  
  147.       
  148.       
  149.     private void fit() {  
  150.         for(int i = 0; i < n_iter; i++)//权重分量调整n_iter次  
  151.         {  
  152.             for(int j = 0; j < exercise.size(); j++)//对于训练集中的每条数据进行训练  
  153.             {  
  154.                 int real_result = exercise.get(j).type;//y  
  155.                 int calculate_result = CalculateResult(exercise.get(j));//y'  
  156.                 double delta0 = eta * (real_result - calculate_result);//计算阈值更新  
  157.                 w0 += delta0;//阈值更新  
  158.                 weights[0] = w0;//更新w[0]  
  159.                 for(int k = 0; k < exercise.get(j).getDimention(); k++)//更新权重向量其它分量  
  160.                 {  
  161.                     double delta = eta * (real_result - calculate_result) * exercise.get(j).vector[k];  
  162.                     //Δw=η*(y-y')*X  
  163.                     weights[k+1] += delta;  
  164.                     //w=w+Δw  
  165.                 }  
  166.                   
  167.             }  
  168.         }  
  169.     }  
  170.   
  171.     private int CalculateResult(Data data) {  
  172.         double z = w0 * x0;  
  173.         for(int i = 0; i < data.dimention; i++)  
  174.             z += data.vector[i] * weights[i+1];  
  175.         //z=w0x0+w1x1+...+WmXm  
  176.         //激活函数  
  177.         if(z>=0)  
  178.             return 1;  
  179.         else  
  180.             return -1;  
  181.     }  
  182.   
  183.     private void showWeigths()  
  184.     {  
  185.         for(double w : weights)  
  186.             System.out.println(w);  
  187.     }  
  188.   
  189.     private void predict(Data data) {  
  190.         int type = CalculateResult(data);  
  191.         if(type == data.getType())  
  192.         {  
  193.             //System.out.println("预测正确");  
  194.         }  
  195.         else  
  196.         {  
  197.             //System.out.println("预测错误");  
  198.             error ++;  
  199.         }  
  200.         testSum ++;  
  201.     }  
  202.   
  203.       
  204. }  


运行结果: 

-0.22000000000000017
-0.4416843982815453
0.442444202054685
总共测试10000条数据,有17条错误,错误率:0.16999999999999998%

 

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