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

R语言:如何高效将矩阵NA值替换为邻域均值?

问题描述

我需要将10×10矩阵中的NA值替换为其邻域的均值,已编写的代码计算效率极低,请问有没有更高效的实现思路或函数?

现有代码

get_neighbor <- function(matrix, x=1,y=1){

  z <- complex(real = rep(1:nrow(matrix), ncol(matrix)),
               imaginary = rep(1:ncol(matrix), each = nrow(matrix)))
  
  lookup <- lapply(seq_along(z), function(x){
    # 计算距离
    dist <- which(abs(z - z[x]) < 2)
    # 移除自身元素
    dist[which(dist != x)]
  })
  index <- (y-1)*(nrow(matrix))+x
  matrix[lookup[[index]]]
  
}

nn_mean <- function(a){
  if(sum(is.na(a))!=ncol(a)*nrow(a)){
    C <- permutations(2, 2, c(1,dim(a)[1]), repeats.allowed = T)
    Borders <- data.frame(matrix(data = 0, ncol = 2, nrow = nrow(a)*2 + ncol(a)*2 - 4))
    Borders[1:nrow(a), 1] <- 1:nrow(a); Borders[1:nrow(a), 2] <- 1
    for(i in 2:(ncol(a)-1)){
      Borders[i + nrow(a) - 1, 2] <- i; Borders[i + 2*(nrow(a) - 1) - 1, 2] <- i
      Borders[i + nrow(a) - 1, 1] <- 1; Borders[i + 2*(nrow(a) - 1) - 1, 1] <- nrow(a)
    }
    Borders[1:ncol(a) + 3*(nrow(a))-4, 2] <- ncol(a)
    Borders[1:ncol(a) + 3*(nrow(a))-4, 1] <- 1:ncol(a)
    
    id <- which(is.na(a), arr.ind = T)
    id <- data.frame(cbind(id, rep(0, nrow(id))))
    
    while(nrow(id)!=0){
      
      for(i in 1:nrow(id)){
        id[i,3] <- sum(is.na(get_neighbor(a, id[i, 1], id[i, 2])))
      }
      
      max_na <- max(id[, 3])
      for(i in 1:(nrow(a)*2 + ncol(a)*2 - 4)){
        if(is.na(a[Borders[i, 1], Borders[i, 2]]) & sum(is.na(get_neighbor(a, Borders[i, 1], Borders[i, 2]))) == 5){
          index <- which(id[,1] == Borders[i, 1] & id[,2] == Borders[i, 2])
          id[index, 3] <- max_na +1
        }
      }
      
      for(i in 1:4){
        if(is.na(a[C[i,1], C[i,2]]) & sum(is.na(get_neighbor(a, C[i, 1], C[i, 2]))) == 3){
          index <- which(id[,1] == C[i, 1] & id[,2] == C[i, 2])
          id[index, 3] <- max_na +1
        }
      }
      
      id <- id[order(id[,3]),]
      index <- which(id[,3]== min(id[,3]))
      for(i in 1:length(index)){
        a[id[i, 1], id[i, 2]] <- mean(get_neighbor(a, id[i, 1], id[i, 2]), na.rm = T)
        if(is.nan(a[id[i, 1], id[i, 2]])){a[id[i, 1], id[i, 2]] <- NA}
      }
      #print(a)
      id <- which(is.na(a), arr.ind = T)
      id <- data.frame(cbind(id, rep(0, nrow(id))))
      
    }
  }
  return(a)
}

示例

a <- matrix(data = runif(100, 0, 10), ncol = 10, nrow = 10)
a[a<2] <- NA  

a
          [,1]     [,2]     [,3]     [,4]     [,5]     [,6]     [,7]     [,8]     [,9]    [,10]
 [1,] 2.313512       NA 5.311104 2.832978 9.917106 2.734799 7.309386       NA 4.794476 6.479147
 [2,] 8.855676 7.555101 8.369477 6.346744 7.727896       NA 9.019421 5.061894 9.116066 6.732293
 [3,] 2.948539 7.440258 6.918414 2.155361 3.511407 5.601253       NA 6.561557 9.543535 4.082592
 [4,] 8.455382 9.169974       NA 4.978224 6.202393       NA 9.435753 9.411371       NA 2.128417
 [5,] 7.744456 3.333072 6.975128 5.876849 4.044768 2.948399 5.067653       NA 6.039412 7.350782
 [6,] 8.793417 9.683755 8.053603 7.406450 6.348171 3.122946 9.378282 5.808363 7.923061 6.415419
 [7,] 4.759612 3.431247 4.123641 6.899569 4.464683 6.588431 5.985248 7.962148 6.668238 4.503556
 [8,] 5.992242       NA 7.099657 6.446650       NA 8.448873 5.884961       NA 2.209453 8.103988
 [9,] 6.383036       NA       NA 5.499157 6.972433 3.129470 3.284383 9.150565 8.484186 4.672878
[10,]       NA       NA 4.258936       NA 9.015525       NA       NA       NA       NA 6.639832

nn_mean(a)

          [,1]     [,2]     [,3]     [,4]     [,5]     [,6]     [,7]     [,8]     [,9]    [,10]
 [1,] 2.313512 6.480974 5.311104 2.832978 9.917106 2.734799 7.309386 7.060248 4.794476 6.479147
 [2,] 8.855676 7.555101 8.369477 6.346744 7.727896 6.545895 9.019421 5.061894 9.116066 6.732293
 [3,] 2.948539 7.440258 6.918414 2.155361 3.511407 5.601253 7.111993 6.561557 9.543535 4.082592
 [4,] 8.455382 9.169974 5.855910 4.978224 6.202393 5.258804 9.435753 9.411371 6.587278 2.128417
 [5,] 7.744456 3.333072 6.975128 5.876849 4.044768 2.948399 5.067653 7.580556 6.039412 7.350782
 [6,] 8.793417 9.683755 8.053603 7.406450 6.348171 3.122946 9.378282 5.808363 7.923061 6.415419
 [7,] 4.759612 3.431247 4.123641 6.899569 4.464683 6.588431 5.985248 7.962148 6.668238 4.503556
 [8,] 5.992242 5.298239 7.099657 6.446650 6.056158 8.448873 5.884961 6.203648 2.209453 8.103988
 [9,] 6.383036 5.902524 5.834195 5.499157 6.972433 3.129470 3.284383 9.150565 8.484186 4.672878
[10,] 6.383036 5.731883 4.258936 6.436513 9.015525 5.600453 5.291218 6.689444 7.236865 6.639832

高效实现思路与方法

原代码效率低的核心原因是大量嵌套循环和重复计算邻域索引,以下是几种更高效的实现方式:

1. 使用滑动窗口工具包(如zoo)

利用现成的滑动窗口函数实现向量化计算,避免手动循环:

library(zoo)

fill_na_neighbor_mean <- function(mat) {
  # 为矩阵添加边界填充,处理边缘元素邻域不足问题
  padded_mat <- cbind(NA, rbind(NA, mat, NA), NA)
  na_pos <- which(is.na(mat), arr.ind = TRUE)
  
  # 填充第一轮可计算的NA
  for (i in 1:nrow(na_pos)) {
    row <- na_pos[i, 1] + 1
    col <- na_pos[i, 2] + 1
    # 提取3×3邻域并排除自身
    neighbor <- padded_mat[(row-1):(row+1), (col-1):(col+1)]
    neighbor[2,2] <- NA
    mat[na_pos[i,1], na_pos[i,2]] <- mean(neighbor, na.rm = TRUE)
    if (is.nan(mat[na_pos[i,1], na_pos[i,2]])) {
      mat[na_pos[i,1], na_pos[i,2]] <- NA
    }
  }
  
  # 循环填充剩余NA(邻域之前为NA、现在已有值的情况)
  while (any(is.na(mat))) {
    na_pos <- which(is.na(mat), arr.ind = TRUE)
    filled <- FALSE
    for (i in 1:nrow(na_pos)) {
      row <- na_pos[i,1]
      col <- na_pos[i,2]
      row_range <- max(1, row-1):min(nrow(mat), row+1)
      col_range <- max(1, col-1):min(ncol(mat), col+1)
      neighbor <- mat[row_range, col_range]
      neighbor[row_range == row, col_range == col] <- NA
      
      mean_val <- mean(neighbor, na.rm = TRUE)
      if (!is.nan(mean_val)) {
        mat[row, col] <- mean_val
        filled <- TRUE
      }
    }
    if (!filled) break
  }
  
  return(mat)
}

2. 向量化预计算邻域索引

提前为每个位置计算邻域索引,避免重复计算,结合向量化操作提升效率:

fill_na_neighbor_mean_vectorized <- function(mat) {
  n_row <- nrow(mat)
  n_col <- ncol(mat)
  
  # 预计算每个位置的有效邻域行/列范围
  neighbor_rows <- lapply(1:n_row, function(r) max(1, r-1):min(n_row, r+1))
  neighbor_cols <- lapply(1:n_col, function(c) max(1, c-1):min(n_col, c+1))
  
  # 生成所有位置的邻域坐标
  indices <- expand.grid(row = 1:n_row, col = 1:n_col)
  indices$neighbors <- mapply(function(r, c) {
    expand.grid(rn = neighbor_rows[[r]], cn = neighbor_cols[[c]]) |>
      subset(!(rn == r & cn == c))
  }, indices$row, indices$col, SIMPLIFY = FALSE)
  
  # 循环填充NA直到无法继续
  while (any(is.na(mat))) {
    na_pos <- which(is.na(mat), arr.ind = TRUE)
    if (nrow(na_pos) == 0) break
    
    # 批量计算NA位置的邻域均值
    means <- apply(na_pos, 1, function(pos) {
      r <- pos[1]
      c <- pos[2]
      neigh <- indices$neighbors[[(c-1)*n_row + r]]
      mean(mat[neigh$rn, neigh$cn], na.rm = TRUE)
    })
    
    # 替换有效均值,跳过邻域全NA的情况
    valid <- !is.nan(means)
    if (sum(valid) == 0) break
    mat[na_pos[valid, 1], na_pos[valid, 2]] <- means[valid]
  }
  
  return(mat)
}

3. 使用Rcpp编写底层代码(极致效率)

对于更大的矩阵,用Rcpp直接操作内存,避免R语言的循环开销:

#include <Rcpp.h>
using namespace Rcpp;

// [[Rcpp::export]]
NumericMatrix fill_na_neighbor_mean_rcpp(NumericMatrix mat) {
  int n_row = mat.nrow();
  int n_col = mat.ncol();
  NumericMatrix res = clone(mat);
  
  bool changed;
  do {
    changed = false;
    // 遍历每个元素
    for (int i = 0; i < n_row; i++) {
      for (int j = 0; j < n_col; j++) {
        if (NumericVector::is_na(res(i,j))) {
          double sum = 0
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 14:06:18