在当下这个时代,各种AI软件崭露头角,完成以前实现不了的各种任务。其背后的神经结构网络也引起了关注,今天我们就来初步认识神经结构网络。

一.感知机

说起神经结构网络,不由得让人联想到生物的神经系统。如果将神经结构网络比作是人脑的神经系统,那么感知机就像是其中的一个个神经元,但在有些情况会更加的复杂,不能严格意义上的相等。先来看看感知机是怎样运行的吧:

感知机是一个线性的模型,只能处理较为简单的二分问题。通过对不同的因素赋予权值,将各个输入与权值的乘积相加得出总值,加上偏移值后经过函数处理得到最终的结果。举个例子,就比如说是大学里绩点的计算可以类比为此,各个学科的成绩通过学分加权计算,加上奖励分后若是大于某一值就可以保研。

那么,看起来那么复杂的神经网络是如何以感知机为基础建立的呢?

说白了,就是以一层层的感知机相互联系组成的,通过一层层的加权计算得出最终的结果。反过来看,想要得到一个固定的答案,可以反过来不断调整权重实现,这是人脑神经不能完成的,也是接下来要实现的实例。

二.实现

接下来要实现的,是一段划分分割线的代码。

import javax.swing.*;
import java.awt.*;
import java.util.Random;

public class Test {
    public static JButton jb;
    public static JFrame jf;
    public static Graphics g;
    public static Perception per;
    public static int count=0;//计算训练的次数
    public static double[][] dot=new double[100][3];
    public static class Perception{
        private double[] weights;//寻找x,y坐标的权重,根据x*weight[0]+y*weight[1]+bias=0寻找出最终的分割线
        private double bias;//偏移量
        private double rate;//学习速率 要寻找合适的速率,并不是越大就越快
        public Perception(int n,double rate){
            this.rate=rate;
            weights=new double[n];
            for (int i = 0; i < n; i++) {
                weights[i]=Math.random();
            }
            bias=Math.random();//先取随机数,再慢慢调整
        }
        public int predict(double[] dot){
            double sum=0;
            for (int i = 0; i < dot.length; i++) {
                sum+=weights[i]*dot[i];//计算权重和
            }
            return func(sum+bias);
        }
        private int func(double sum){
            return sum>=0?1:-1;//根据x*weight[0]+y*weight[1]+bias=0,此处的sum视为距离所求直线的距离,根据距离返回不同的点的类型
        }
        public void train(double[] dot){
            int prediction=predict(new double[]{dot[0],dot[1]});
            if(prediction!=dot[2]){
                for (int i = 0; i < weights.length; i++) {
                    weights[i]+=rate*(dot[2]-prediction)*dot[i];//根据实际与预测调整权重,同时越大的坐标值需要更大的权重修正误差,乘上坐标调整
                }
                bias+=rate*(dot[2]-prediction);//根据实际与预测的差距调整偏移量
            }//rate只是一个系数,设置为多少都可以,不影响式子表达
        }
    }
    public static void getRandomDot(){
        Random r=new Random();
        for(int i=0;i<100;i++){
            dot[i][0]=r.nextDouble()*800;
            dot[i][1]=r.nextDouble()*600;
            dot[i][2]=(dot[i][0]+dot[i][1]>700)?1:-1;//因为感知机是线性的,定义一条分割线区分两点便于分割
        }
    }

    public static void main(String[] args) {
        per=new Perception(2,0.2);
        getRandomDot();
        showUI();
    }
    public static void showUI(){
        jf=new JFrame("感知机可视化");
        jf.setSize(800,600);
        jf.setDefaultCloseOperation(JFrame.EXIT_ON_CLOSE);
        jf.setLayout(new FlowLayout());
        jf.setLocationRelativeTo(null);
        jb=new JButton("开始训练");
        jf.add(jb);
        jf.setVisible(true);
        g=jf.getGraphics();
        jb.addActionListener(e -> {
            count++;
            g.clearRect(0,0,800,600);
            trainPer();
            devide(g);
            jb.setText("训练 "+count+"次");
            draw(g);
        });
    }
    public static void draw(Graphics g){
        for(int i=0;i<100;i++){
            int x=(int)dot[i][0];
            int y=(int)dot[i][1];
            if(dot[i][2]==1){//画正字
                g.setColor(Color.RED);
                g.drawLine(x-3,y-3,x+3,y-3);
                g.drawLine(x,y+3,x,y-3);
                g.drawLine(x,y,x+3,y);
                g.drawLine(x+3,y+3,x-3,y+3);
                g.drawLine(x-2,y,x-2,y+3);
            }else{//画恶字
                g.setColor(Color.BLACK);
                g.drawLine(x-4,y-3,x+4,y-3);
                g.drawLine(x-1,y-3,x-1,y);
                g.drawLine(x+1,y-3,x+1,y);
                g.drawLine(x-3,y-2,x-1,y-1);
                g.drawLine(x+3,y-2,x+1,y-1);
                g.drawLine(x-3,y,x+3,y);
                g.drawLine(x-3,y+1,x-4,y+3);
                g.drawLine(x-2,y+1,x-2,y+3);
                g.drawLine(x-2,y+3,x+2,y+3);
                g.drawLine(x+2,y+3,x+2,y+1);
                g.fillOval(x-1,y,2,2);
                g.drawLine(x+3,y+1,x+4,y+3);
            }
        }
    }
    private static void devide(Graphics g){
        double[] weights=per.weights;
        double bias=per.bias;
        g.setColor(Color.CYAN);
        g.drawLine(0,(int)(-(weights[0]*0+bias)/weights[1]),800,(int)(-(weights[0]*800+bias)/weights[1]));//划分割线,设计这样两个x坐标能贯穿始末,根据分割线公式求得y坐标
    }
    private static void trainPer(){
        for (int epoch = 0; epoch < 1000; epoch++) {//设定训练次数
            for(double[] data:dot){
                per.train(data);
            }
        }
    }
}

神经网络结构并没有想象中的那么复杂,理解了感知机算是理解了神经结构网络,只要有耐心人人都可以手写一个神经网络。

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐