| 64 | |
| 65 | |
| 66 | @Override |
| 67 | public <Type extends DataSet> double[] getImportanceStats(TreeLearner model, DataSet<Type> data) |
| 68 | { |
| 69 | double[] features = new double[data.getNumFeatures()]; |
| 70 | |
| 71 | if(!(data instanceof ClassificationDataSet)) |
| 72 | throw new RuntimeException("MDI currently only supports classification datasets"); |
| 73 | |
| 74 | List<DataPointPair<Integer>> allData = ((ClassificationDataSet)data).getAsDPPList(); |
| 75 | final int K = ((ClassificationDataSet)data).getClassSize(); |
| 76 | ImpurityScore score = new ImpurityScore(K, im); |
| 77 | for(DataPointPair<Integer> d : allData) |
| 78 | score.addPoint(d.getDataPoint(), d.getPair()); |
| 79 | |
| 80 | visit(model.getTreeNodeVisitor(), score, allData, features, score.getSumOfWeights(), K); |
| 81 | |
| 82 | return features; |
| 83 | } |
| 84 | |
| 85 | private void visit(TreeNodeVisitor node, ImpurityScore score, List<DataPointPair<Integer>> data, final double[] features , final double N, final int K) |
| 86 | { |