【问题标题】:Cross validation and ROC curve using Matlab: how plot mean ROC curve?使用 Matlab 进行交叉验证和 ROC 曲线:如何绘制平均 ROC 曲线?
【发布时间】:2021-01-30 15:31:26
【问题描述】:

我正在使用 k = 10 的 k 折交叉验证。因此,我有 10 条 ROC 曲线。 我想在曲线之间进行平均。我不能只平均 Y 轴上的值(使用 perfcurve),因为返回的向量大小不一样。

[X1,Y1,T1,AUC1] = perfcurve(t_test(1),resp(1),1);
.
.
.
[X10,Y10,T10,AUC10] = perfcurve(t_test(10),resp(10),1);

如何解决这个问题?如何绘制 10 条 ROC 曲线的平均曲线?

【问题讨论】:

    标签: matlab cross-validation roc auc k-fold


    【解决方案1】:

    所以,您有具有不同点数的 k 曲线,在两个维度上都绑定在 [0..1] 区间内。首先,您需要计算指定查询点的每条曲线的插值。现在你有了新的曲线,点数固定,可以计算它们的平均值。 interp1 函数将执行插值部分。

    %% generating sample data
    k = 10;
    X = cell(k, 1);
    Y = cell(k, 1);
    hold on;
    for i=1:k
        n = 10+randi(10);
        X{i} = sort([0 1 rand(1, n)]);
        Y{i} = sort([0 1 rand(1, n)].^.5);
    end
    
    %% Calculating interpolations
    % location of query points
    X2 = linspace(0, 1, 50);
    n = numel(X2);
    % initializing values for different curves at different query points
    Y2 = zeros(k, n);
    for i=1:k
        % finding interpolated values for i-th curve
        Y2(i, :) = interp1(X{i}, Y{i}, X2);
    end
    % finding the mean
    meanY = mean(Y2, 1);
    
    

    请注意,不同的插值方法会影响您的结果。例如,ROC 图数据是一种楼梯数据。要在此类曲线上找到准确的值,您应该使用先前邻居插值方法,而不是 interp1 的默认方法线性插值:

    Y2(i, :) = interp1(X{i}, Y{i}, X2); % linear
    Y3(i, :) = interp1(X{i}, Y{i}, X2, 'previous');
    

    这就是它对最终结果的影响:

    【讨论】:

    • 感谢您的回答。我找到了一个使用 perfcurve 的解决方案,我相信这是一个更简单的任务,但你的答案也很多。
    【解决方案2】:

    我使用 Matlab 的 perfcurve 解决了这个问题。为此,我必须将 "label" 和 "scores" 的向量列表(大小向量 1xn)作为参数传递。因此,perfcurve 函数已经理解为使用 k 倍进行的一组分辨率,并返回平均 ROC 曲线及其置信区间,以及 AUC 及其置信区间。

    [X1,Y1,T1,AUC1] = perfcurve(t_test_list,resp_list,1);

    t_testresp 它们是大小为 1xk 的列表(k 是折叠数/k 折叠),列表的每个元素都是一个 1xn 向量,带有分数和标签。

    resp = nnet(x_test(i));
    t_test_act = t_test(i); 
    

    resp 具有 2xn 格式(n 是预测样本的数量)。有两个类。

    t_test_act 包含当前测试集的标签,已形成2xn,由0和1组成(每列有1和0,表示样本的真实类别) .

    resp_list{i} = resp(1,:)  %(scores)
    t_test_list{i} = t_test_act(1,:) %(labels)
    [X1,Y1,T1,AUC1] = perfcurve(t_test_list,resp_list,1);
    

    【讨论】:

      猜你喜欢
      • 2021-05-08
      • 2012-09-11
      • 2019-02-27
      • 2019-08-15
      • 2018-01-15
      • 2020-01-02
      • 2018-12-28
      • 2019-01-25
      • 1970-01-01
      相关资源
      最近更新 更多