You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义kNN缺失值填充性能优化求助:运行超12小时亟待提速

kNN算法优化求助:考试数据缺失得分填充性能瓶颈解决

问题背景

数据集为选择题考试数据,字段包括:student_id(学生标识)、course(课程)、year(修课年份)、question_id(试题标识)、ans(是否作答)、score(题目得分)。需使用k=5的kNN算法填充未答题的缺失得分。

当前实现思路

  • 按年份和课程过滤数据集(保留至少作答过一道题的学生)
  • 对每个有未答题的学生,遍历其所有未答题:
    • 针对每道未答题,筛选出所有作答过该题的其他学生
    • 计算自定义距离:同一试题得分相同则距离为0,不同则为1,总距离为所有试题的该值之和
    • 取距离最小的5名学生的该题得分均值作为填充值

性能现状

数据集共40万行(试题-学生记录),涵盖5个年份、3门课程、约6000名唯一学生,未答题率20%。当前基于data.table的R语言实现运行时长超过12小时,希望压缩至分钟级或数小时。

瓶颈分析

核心瓶颈是重复计算学生间的距离,尝试预计算距离矩阵未成功,同时不清楚R内置knn函数是否支持自定义距离函数。

现有代码

距离计算函数

compute_distance <- function(student_a, student_b, data_course, current_question) {
  # Get scores for all other questions except the current one
  scores_a <- data_course[student_id == student_a & question_id != current_question, SCORE]
  scores_b <- data_course[student_id == student_b & question_id != current_question, SCORE]
  
  # Ensure both score sets have the same questions
  if (length(scores_a) == length(scores_b)) {
    # Binary distance (0 if same score, 1 if different)
    dist <- sum(scores_a != scores_b)
    return(dist)
  } else {
    return(Inf)  # If questions don't match, return a large distance
  }
}

主填充函数

knn_impute <- function(data, k = 5) {
  
  # Start the entire process timer
  total_start_time <- Sys.time()
  
  # Iterate over each unique year
  for (year_x in unique(data$year)) {
    # Track the time for each year
    year_start_time <- Sys.time()
    
    # Print current year
    cat("Processing year:", year_x, "\n")
    
    # Filter data for this year
    data_year <- data[year == year_x]
    
    # Iterate over each unique course within the year
    for (course_x in unique(data_year$course)) {
      # Track the time for each course
      course_start_time <- Sys.time()
      
      # Print current course
      cat("  Processing course:", course_x, "\n")
      
      # Filter data for this course within the year
      data_course <- data_year[course == course_x]
      
      # Get unique unanswered students for this year and course
      unanswered_students <- unique(data_course[ans == 0, student_id])
      
      # Iterate over each student who has unanswered questions
      for (student_x in unanswered_students) {
        # Get all unanswered questions for student_x
        unanswered_questions <- data_course[student_id == student_x & ans == 0, question_id]
        
        # Iterate over each unanswered question for student_x
        for (question_x in unanswered_questions) {
          
          # Get all students who answered the current question
          answered_students <- data_course[question_id == question_x & ans != 0, student_id]
          
          # If no one answered the question, skip it
          if (length(answered_students) == 0) next
          
          # Calculate distances and store directly in a vector for sorting
          distances <- sapply(answered_students, function(student_y) {
            compute_distance(student_x, student_y, data_course, question_x)
          })
          
          # Sort distances and select the k nearest neighbors
          nearest_neighbors <- answered_students[order(distances)[1:min(k, length(distances))]]
          
          # Impute the score for student_x's unanswered question
          imputed_score <- mean(data_course[student_id %in% nearest_neighbors & question_id == question_x, score], na.rm = TRUE)
          
          # Update the score for student_x in the data
          data[student_id == student_x & question_id == question_x, score := imputed_score]
        }
      }
      
      # Calculate and print course duration
      course_end_time <- Sys.time()
      course_duration <- as.numeric(difftime(course_end_time, course_start_time, units = "secs"))
      cat("  Time spent on course:", course_x, "=", course_duration, "seconds\n")
    }
    
    # Calculate and print year duration
    year_end_time <- Sys.time()
    year_duration <- as.numeric(difftime(year_end_time, year_start_time, units = "secs"))
    cat("Time spent on year:", year_x, "=", year_duration, "seconds\n\n")
  }
  
  # Print total time spent
  total_end_time <- Sys.time()
  total_duration <- as.numeric(difftime(total_end_time, total_start_time, units = "secs"))
  cat("Total time spent on kNN imputation =", total_duration, "seconds\n")
  
  return(data)
}

# Call your function and run the imputation
imputed_data <- knn_impute(filtered_data, k = 5)

示例数据

student_idquestion_idcourseyearansscore
110Course_A201913
215Course_B202000
320Course_A202115
425Course_C201914
530Course_A202000

求助需求

恳请提供性能优化建议或更好的解决方案,将运行耗时压缩至分钟级或数小时。


内容的提问来源于stack exchange,提问作者jacopoburelli

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 01:44:56