代码之家  ›  专栏  ›  技术社区  ›  Christopher

段落向量模型的交叉验证

  •  0
  • Christopher  · 技术社区  · 7 年前

    我在尝试对段落向量模型应用交叉验证时遇到了一个错误:

    import numpy as np
    import pandas as pd
    from sklearn.linear_model import LogisticRegression
    from sklearn.model_selection import cross_val_score
    from sklearn.pipeline import Pipeline
    from gensim.sklearn_api import D2VTransformer
    
    data = pd.read_csv('https://pastebin.com/raw/bSGWiBfs')
    np.random.seed(0)
    
    X_train = data.apply(lambda r: simple_preprocess(r['text'], min_len=2), axis=1)
    y_train = data.label
    
    model = D2VTransformer(size=10, min_count=1, iter=5, seed=1)
    clf = LogisticRegression(random_state=0)
    
    pipeline = Pipeline([
            ('vec', model),
            ('clf', clf)
        ])
    
    pipeline.fit(X_train, y_train)
    
    score = pipeline.score(X_train, y_train)
    print("Score:", score) # This works
    cval = cross_val_score(pipeline, X_train, y_train, scoring='accuracy', cv=3)
    print("Cross-Validation:", cval) # This doesn't work
    

    关键错误:0

    X_train 在里面 cross_val_score model.transform(X_train) model.fit_transform(X_train) . 此外,我还对原始输入数据进行了同样的尝试( data.text ),而不是预处理的文本。我怀疑这本书的格式一定有问题 火车 对于交叉验证,与 .score 用于管道的函数,工作正常。我还注意到 交叉评分 合作 CountVectorizer()

    有人发现错误了吗?

    1 回复  |  直到 7 年前
        1
  •  1
  •   Vivek Kumar    7 年前

    不,这与从 model . 它与 cross_val_score .

    交叉评分 将根据 cv 帕拉姆。为此,它将执行以下操作:

    for train, test in splitter.split(X_train, y_train):
        new_X_train, new_y_train = X_train[train], y_train[train]
    

    X_train 是一个 pandas.Series 对象,在该对象中基于索引的选择不是这样工作的。见此: https://pandas.pydata.org/pandas-docs/stable/indexing.html#selection-by-position

    更改此行:

    X_train = data.apply(lambda r: simple_preprocess(r['text'], min_len=2), axis=1)
    

    致:

    # Access the internal numpy array
    X_train = data.apply(lambda r: simple_preprocess(r['text'], min_len=2), axis=1).values
    
    OR
    
    # Convert series to list
    X_train = data.apply(lambda r: simple_preprocess(r['text'], min_len=2), axis=1).tolist()
    
    推荐文章