数据集划分——DUPLEX


引言

在机器学习中数据集划分的方式一般分为以下几类

? Simple random sampling (SRS)

? Trial-and-error methods

? Systematic sampling

? Convenience sampling

? CADEX, DUPLEX

? Stratified sampling

DUPLEX划分

其中DUPLEX划分是CADEX方法的扩展实现,根据样本之间的欧氏距离来选择样本。具体来说就是从数据集T的两个最远的样本开始,然后重复地选择与之前采样的样本有最大距离的样本。该方法保证了了对数据集T的最大覆盖。

缺点:计算复杂性使DUPLEX无法用于大型高维数据集。

具体步骤如下

 代码实现

从X的样本矩阵(一行代表一个样本),挑选K个作为测试集,(K大于样本数的一半,则会返回训练集)

import numpy as np

def duplex(X, k, progress=None, isCancelled=None):
    n = X.shape[0]
    p = X.shape[1]
    if n/2 < k:
        k = n-k
        outReverse = True
    else:
        outReverse = False

    if k==0 or k==1: # FIXME 需要处理k<2的情况
        if outReverse:
            return list(range(n))
        else:
            return []

    X1 = np.append(X, np.arange(n)[:, np.newaxis], axis=1)
    invCov = np.eye(p)

    model = list()
    test = list()
    rest = list(range(n))

    p = np.tril(_pdist(X1[rest, :-1], X1[rest, :-1], invCov))
    f = np.where(p == p.max().max())
    i1 = int(X1[rest, :][f[0][0], -1])
    model.append(i1)
    rest.remove(i1)
    i2 = int(X1[rest, :][f[1][0], -1])
    model.append(i2)
    rest.remove(i2)

    p = np.tril(_pdist(X1[rest, :-1], X1[rest, :-1], invCov))
    f = np.where(p == p.max().max())
    i1 = int(X1[rest, :][f[0][0], -1])
    test.append(i1)
    rest.remove(i1)
    i2 = int(X1[rest, :][f[1][0], -1])
    test.append(i2)
    rest.remove(i2)

    if callable(progress):
        progress(100*2/k)
    if callable(isCancelled):
        if isCancelled():
            return test

    for j in range(k-2):
        p = _pdist(X1[model, :-1], X1[rest, :-1], invCov)
        i = np.argmax(p.min(axis=0))
        ii = int(X1[rest, :][i, -1])
        model.append(ii)
        rest.remove(ii)

        p = _pdist(X1[test, :-1], X1[rest, :-1], invCov)
        i = np.argmax(p.min(axis=0))
        ii = int(X1[rest, :][i, -1])
        test.append(ii)
        rest.remove(ii)

        # 进度条以及中断
        if callable(progress):
            progress(100*j/k)
        if callable(isCancelled):
            if isCancelled():
                break
    if outReverse:
        return list(set(range(n)).difference(test))
    else:
        return test

References

[1] Reitermanova Z. Data splitting[C]//WDS. 2010, 10: 31-36.