Skip to content

Commit 65d0fe9

Browse files
committed
array2dt添加测试,放心使用
1 parent 872c400 commit 65d0fe9

3 files changed

Lines changed: 124 additions & 18 deletions

File tree

NAMESPACE

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -240,6 +240,7 @@ importFrom(crayon,red)
240240
importFrom(crayon,underline)
241241
importFrom(data.table,CJ)
242242
importFrom(data.table,as.data.table)
243+
importFrom(data.table,chmatch)
243244
importFrom(data.table,data.table)
244245
importFrom(data.table,dcast)
245246
importFrom(data.table,fread)

R/tools_data.table.R

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ fwrite2 <- function(x, file) {
3333
#'
3434
#' @example R/examples/array2dt.R
3535
#'
36-
#' @importFrom data.table CJ
36+
#' @importFrom data.table CJ chmatch
3737
#' @export
3838
array2dt <- function(arr, dimnames) {
3939
# 注意顺序要匹配,aperm非常必要!
@@ -47,12 +47,15 @@ dt2array <- function(dt, value_col = "value") {
4747
# 获取维度列(除了值列之外的所有列)
4848
dim_cols <- setdiff(names(dt), value_col)
4949

50-
# 获取每个维度的唯一值(保持原顺序)
51-
dimnames <- lapply(dt[, ..dim_cols], function(x) unique(x))
52-
53-
# 创建 array 并用 aperm 调整维度顺序
54-
dims <- lengths(dimnames) # 计算维度大小
55-
dim = rev(dims) %>% setNames(NULL)
56-
arr <- array(dt[[value_col]], dim = dim, dimnames = rev(dimnames))
57-
aperm(arr)
50+
# 获取每个维度的唯一值(保持首次出现的顺序)
51+
dimnames <- lapply(dt[, ..dim_cols], unique)
52+
dims <- setNames(lengths(dimnames), NULL)
53+
54+
# 按维度取值定位填充,不依赖 dt 的行顺序;缺失组合记为 NA
55+
# 字符列用 data.table::chmatch (比 match 快约 5 倍),其余用 match
56+
.match <- function(x, levs) if (is.character(x)) chmatch(x, levs) else match(x, levs)
57+
idx <- mapply(function(col, levs) .match(dt[[col]], levs), dim_cols, dimnames)
58+
arr <- array(dt[[value_col]][NA_integer_], dim = dims, dimnames = dimnames)
59+
arr[idx] <- dt[[value_col]]
60+
arr
5861
}

tests/testthat/test-array2dt.R

Lines changed: 111 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,19 +1,121 @@
1-
test_that("array2dt works", {
2-
# 测试往返转换
3-
arr <- array(1:6,
1+
## helpers ---------------------------------------------------------------------
2+
make_arr <- function() {
3+
array(1:6,
44
dim = c(2, 3),
5+
dimnames = list(site = c("A", "B"), date = c("d1", "d2", "d3"))
6+
)
7+
}
8+
9+
make_arr3 <- function() {
10+
array(1:24,
11+
dim = c(2, 3, 4),
512
dimnames = list(
613
site = c("A", "B"),
7-
date = c("d1", "d2", "d3")
8-
# var = c("v1", "v2", "v3", "v4")
14+
date = c("d1", "d2", "d3"),
15+
var = paste0("v", 1:4)
916
)
1017
)
18+
}
1119

12-
# array -> dt
20+
## array2dt --------------------------------------------------------------------
21+
test_that("array2dt: 维度列、行数与列名正确", {
22+
arr <- make_arr()
1323
dt <- array2dt(arr, dimnames(arr))
24+
25+
expect_s3_class(dt, "data.table")
26+
expect_equal(nrow(dt), prod(dim(arr)))
27+
expect_equal(names(dt), c("site", "date", "value"))
28+
})
29+
30+
test_that("array2dt: 展开顺序为最后一维变化最快 (CJ 顺序)", {
31+
arr <- make_arr()
32+
dt <- array2dt(arr, dimnames(arr))
33+
34+
# 第一维 (site) 变化最慢,最后一维 (date) 变化最快
35+
expect_equal(dt$site, rep(c("A", "B"), each = 3))
36+
expect_equal(dt$date, rep(c("d1", "d2", "d3"), times = 2))
37+
# 值需与 array 下标严格对应
1438
expect_equal(dt[site == "A", value], c(1, 3, 5))
39+
expect_equal(dt[site == "B", value], c(2, 4, 6))
40+
})
41+
42+
test_that("array2dt: 支持三维数组", {
43+
arr <- make_arr3()
44+
dt <- array2dt(arr, dimnames(arr))
45+
46+
expect_equal(nrow(dt), 24)
47+
expect_equal(names(dt), c("site", "date", "var", "value"))
48+
expect_equal(dt[site == "A" & date == "d1", value], as.vector(arr["A", "d1", ]))
49+
})
50+
51+
## dt2array --------------------------------------------------------------------
52+
test_that("dt2array: 二维往返转换", {
53+
arr <- make_arr()
54+
expect_equal(dt2array(array2dt(arr, dimnames(arr))), arr)
55+
})
56+
57+
test_that("dt2array: 三维往返转换", {
58+
arr <- make_arr3()
59+
expect_equal(dt2array(array2dt(arr, dimnames(arr))), arr)
60+
})
61+
62+
test_that("dt2array: 值的填充不依赖 dt 的行顺序 (回归测试)", {
63+
arr <- make_arr3()
64+
dt <- array2dt(arr, dimnames(arr))
65+
66+
# 维度 level 顺序恰好保持时, 应与原数组完全一致
67+
expect_equal(dt2array(dt[order(value)]), arr)
68+
expect_equal(dt2array(data.table::setkey(data.table::copy(dt), var)), arr)
69+
70+
# 完全打乱行序: level 顺序按首次出现推断, 重排回原顺序后值应一致
71+
set.seed(1)
72+
r <- dt2array(dt[sample(.N)])
73+
r <- r[dimnames(arr)$site, dimnames(arr)$date, dimnames(arr)$var]
74+
expect_equal(r, arr)
75+
})
76+
77+
test_that("dt2array: 缺失的维度组合记为 NA", {
78+
arr <- make_arr()
79+
dt <- array2dt(arr, dimnames(arr))
80+
dt_miss <- dt[!(site == "A" & date == "d2")]
81+
82+
r <- dt2array(dt_miss)
83+
expect_equal(dim(r), dim(arr)) # 维度不缩减 (其余行仍含全部 level)
84+
expect_true(is.na(r["A", "d2"]))
85+
expect_equal(sum(is.na(r)), 1L)
86+
})
87+
88+
test_that("dt2array: 保留值的数据类型", {
89+
# double
90+
arr_d <- array(c(1.5, 2.5, 3.5, 4.5, 5.5, 6.5),
91+
dim = c(2, 3), dimnames = list(s = c("A", "B"), d = c("d1", "d2", "d3")))
92+
r_d <- dt2array(array2dt(arr_d, dimnames(arr_d)))
93+
expect_type(r_d, "double")
94+
expect_equal(r_d, arr_d)
95+
96+
# character
97+
arr_c <- array(letters[1:6],
98+
dim = c(2, 3), dimnames = list(s = c("A", "B"), d = c("d1", "d2", "d3")))
99+
r_c <- dt2array(array2dt(arr_c, dimnames(arr_c)))
100+
expect_type(r_c, "character")
101+
expect_equal(r_c, arr_c)
102+
})
103+
104+
test_that("dt2array: 支持非字符维度列 (走 match 分支)", {
105+
arr <- array(1:6, dim = c(2, 3),
106+
dimnames = list(id = c(10L, 20L), date = c("d1", "d2", "d3")))
107+
dt <- array2dt(arr, dimnames(arr))
108+
dt[, id := as.integer(id)] # id 为 integer 列, 由 match 处理
109+
110+
set.seed(1)
111+
r <- dt2array(dt[sample(.N)])
112+
expect_equal(r[as.character(c(10, 20)), dimnames(arr)$date], arr)
113+
})
114+
115+
test_that("dt2array: 支持自定义值列名", {
116+
arr <- make_arr()
117+
dt <- array2dt(arr, dimnames(arr))
118+
data.table::setnames(dt, "value", "val")
15119

16-
# dt -> array
17-
arr2 <- dt2array(dt)
18-
expect_equal(arr, arr2)
120+
expect_equal(dt2array(dt, value_col = "val"), arr)
19121
})

0 commit comments

Comments
 (0)