【问题标题】:libsvm java implementationlibsvm java实现
【发布时间】:2012-06-03 06:55:45
【问题描述】:

我正在尝试为 libsvm 使用 java 绑定:

http://www.csie.ntu.edu.tw/~cjlin/libsvm/

我已经实现了一个“简单”的例子,它很容易在 y 中线性分离。数据定义为:

double[][] train = new double[1000][]; 
double[][] test = new double[10][];

for (int i = 0; i < train.length; i++){
    if (i+1 > (train.length/2)){        // 50% positive
        double[] vals = {1,0,i+i};
        train[i] = vals;
    } else {
        double[] vals = {0,0,i-i-i-2}; // 50% negative
        train[i] = vals;
    }           
}

第一个“特征”是类,训练集的定义类似。

训练模型:

private svm_model svmTrain() {
    svm_problem prob = new svm_problem();
    int dataCount = train.length;
    prob.y = new double[dataCount];
    prob.l = dataCount;
    prob.x = new svm_node[dataCount][];     

    for (int i = 0; i < dataCount; i++){            
        double[] features = train[i];
        prob.x[i] = new svm_node[features.length-1];
        for (int j = 1; j < features.length; j++){
            svm_node node = new svm_node();
            node.index = j;
            node.value = features[j];
            prob.x[i][j-1] = node;
        }           
        prob.y[i] = features[0];
    }               

    svm_parameter param = new svm_parameter();
    param.probability = 1;
    param.gamma = 0.5;
    param.nu = 0.5;
    param.C = 1;
    param.svm_type = svm_parameter.C_SVC;
    param.kernel_type = svm_parameter.LINEAR;       
    param.cache_size = 20000;
    param.eps = 0.001;      

    svm_model model = svm.svm_train(prob, param);

    return model;
}

然后评估我使用的模型:

public int evaluate(double[] features) {
    svm_node node = new svm_node();
    for (int i = 1; i < features.length; i++){
        node.index = i;
        node.value = features[i];
    }
    svm_node[] nodes = new svm_node[1];
    nodes[0] = node;

    int totalClasses = 2;       
    int[] labels = new int[totalClasses];
    svm.svm_get_labels(_model,labels);

    double[] prob_estimates = new double[totalClasses];
    double v = svm.svm_predict_probability(_model, nodes, prob_estimates);

    for (int i = 0; i < totalClasses; i++){
        System.out.print("(" + labels[i] + ":" + prob_estimates[i] + ")");
    }
    System.out.println("(Actual:" + features[0] + " Prediction:" + v + ")");            

    return (int)v;
}

其中传递的数组是测试集中的一个点。

结果总是返回类 0。 确切的结果是:

(0:0.9882998314585194)(1:0.011700168541480586)(Actual:0.0 Prediction:0.0)
(0:0.9883952943701599)(1:0.011604705629839989)(Actual:0.0 Prediction:0.0)
(0:0.9884899803606306)(1:0.011510019639369528)(Actual:0.0 Prediction:0.0)
(0:0.9885838957058696)(1:0.011416104294130458)(Actual:0.0 Prediction:0.0)
(0:0.9886770466322342)(1:0.011322953367765776)(Actual:0.0 Prediction:0.0)
(0:0.9870913229268679)(1:0.012908677073132284)(Actual:1.0 Prediction:0.0)
(0:0.9868781382588805)(1:0.013121861741119505)(Actual:1.0 Prediction:0.0)
(0:0.986661444476744)(1:0.013338555523255982)(Actual:1.0 Prediction:0.0)
(0:0.9864411843906802)(1:0.013558815609319848)(Actual:1.0 Prediction:0.0)
(0:0.9862172999068877)(1:0.013782700093112332)(Actual:1.0 Prediction:0.0)

谁能解释为什么这个分类器不起作用? 是否有步骤我搞砸了,或者我缺少了什么步骤?

谢谢

【问题讨论】:

    标签: java svm libsvm


    【解决方案1】:

    在我看来,您的评估方法是错误的。应该是这样的:

    public double evaluate(double[] features, svm_model model) 
    {
        svm_node[] nodes = new svm_node[features.length-1];
        for (int i = 1; i < features.length; i++)
        {
            svm_node node = new svm_node();
            node.index = i;
            node.value = features[i];
    
            nodes[i-1] = node;
        }
    
        int totalClasses = 2;       
        int[] labels = new int[totalClasses];
        svm.svm_get_labels(model,labels);
    
        double[] prob_estimates = new double[totalClasses];
        double v = svm.svm_predict_probability(model, nodes, prob_estimates);
    
        for (int i = 0; i < totalClasses; i++){
            System.out.print("(" + labels[i] + ":" + prob_estimates[i] + ")");
        }
        System.out.println("(Actual:" + features[0] + " Prediction:" + v + ")");            
    
        return v;
    }
    

    【讨论】:

    • 你能解释一下问题代码中的错误是什么吗?我在发现错误时遇到问题! :(
    【解决方案2】:

    这是我使用以下 R 代码中的数据测试过的上述示例的返工:http://cbio.ensmp.fr/~jvert/svn/tutorials/practical/svmbasic/svmbasic_notes.pdf

    import libsvm.*;
    
    public class libsvmTest {
    
      public static void main(String [] args) {
    
          double[][] xtrain = ...
          double[][] xtest = ...
          double[][] ytrain = ...
          double[][] ytest = ...
    
          svm_model m = svmTrain(xtrain,ytrain);
    
          double[] ypred = svmPredict(xtest, m); 
    
          for (int i = 0; i < xtest.length; i++){
              System.out.println("(Actual:" + ytest[i][0] + " Prediction:" + ypred[i] + ")"); 
          }  
    
      }
    
      static svm_model svmTrain(double[][] xtrain, double[][] ytrain) {
            svm_problem prob = new svm_problem();
            int recordCount = xtrain.length;
            int featureCount = xtrain[0].length;
            prob.y = new double[recordCount];
            prob.l = recordCount;
            prob.x = new svm_node[recordCount][featureCount];     
    
            for (int i = 0; i < recordCount; i++){            
                double[] features = xtrain[i];
                prob.x[i] = new svm_node[features.length];
                for (int j = 0; j < features.length; j++){
                    svm_node node = new svm_node();
                    node.index = j;
                    node.value = features[j];
                    prob.x[i][j] = node;
                }           
                prob.y[i] = ytrain[i][0];
            }               
    
            svm_parameter param = new svm_parameter();
            param.probability = 1;
            param.gamma = 0.5;
            param.nu = 0.5;
            param.C = 100;
            param.svm_type = svm_parameter.C_SVC;
            param.kernel_type = svm_parameter.LINEAR;       
            param.cache_size = 20000;
            param.eps = 0.001;      
    
            svm_model model = svm.svm_train(prob, param);
    
            return model;
        }  
    
      static double[] svmPredict(double[][] xtest, svm_model model) 
      {
    
          double[] yPred = new double[xtest.length];
    
          for(int k = 0; k < xtest.length; k++){
    
            double[] fVector = xtest[k];
    
            svm_node[] nodes = new svm_node[fVector.length];
            for (int i = 0; i < fVector.length; i++)
            {
                svm_node node = new svm_node();
                node.index = i;
                node.value = fVector[i];
                nodes[i] = node;
            }
    
            int totalClasses = 2;       
            int[] labels = new int[totalClasses];
            svm.svm_get_labels(model,labels);
    
            double[] prob_estimates = new double[totalClasses];
            yPred[k] = svm.svm_predict_probability(model, nodes, prob_estimates);
    
          }
    
          return yPred;
      } 
    
    
    }
    

    这是输出:

    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:1.0 Prediction:1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    (Actual:-1.0 Prediction:-1.0)
    

    【讨论】:

    • 非常感谢您提供有用的代码。为什么使用 param.probability = 1;?其次,你知道如果一个人的班级不平衡,如何设置权重吗?我的意思是 C 参数的权重。
    • 调用 svm.svm_predict_probability() 时 prob_estimates 不会失去作用域吗?
    • 这只是一个帮助开始使用 LIBSVM 的帖子;从那里,用户可以根据问题确定什么工作。对于这方面的问题,我建议你访问这个包的维护者的网站:csie.ntu.edu.tw/~cjlin/libsvm/faq.html#/…
    【解决方案3】:

    我对 LibSVM 的 java 实现做了一个稍微重构的版本,您可能会发现它更易于使用: https://github.com/syeedibnfaiz/libsvm-java-kernel。 查看 Demo.java 类以了解如何使用它。

    【讨论】:

      猜你喜欢
      • 2012-11-11
      • 2013-04-19
      • 2012-04-20
      • 2013-12-25
      • 2012-02-04
      • 2013-02-20
      • 2017-10-12
      • 2014-10-21
      • 2012-05-24
      相关资源
      最近更新 更多