8.1 concatenate家族与split 本节摘要:concatenate 沿已有轴拼接(结果维度不变),stack 沿新轴堆叠(维度加一),两者必然拷贝——拼接后的排列无法用一套 strides 描述。split 家族反向切分,纯规则切分给视图,不均等切分用 arraysplit。本节给出轴选择决策、hstack 与 vstack 的坑,以及"循环 append 陷阱"的修复。 合并为什么必然拷贝 两块独立内存拼成一块,中间不可能有"跳转指令"——ndarray 的连续性要求数据块整体可寻址,所以合并的结果必然是一块全新分配、逐元素搬运的内存: 代价意识:合并 2 份 100MB 的数组要新分配 200MB 并搬运,峰值瞬间三份并存。
本节摘要:concatenate 沿已有轴拼接(结果维度不变),stack 沿新轴堆叠(维度加一),两者必然拷贝——拼接后的排列无法用一套 strides 描述。split 家族反向切分,纯规则切分给视图,不均等切分用 array_split。本节给出轴选择决策、hstack 与 vstack 的坑,以及"循环 append 陷阱"的修复。
两块独立内存拼成一块,中间不可能有"跳转指令"——ndarray 的连续性要求数据块整体可寻址,所以合并的结果必然是一块全新分配、逐元素搬运的内存:
import numpy as np a = np.arange(6).reshape(2, 3) b = np.arange(6, 12).reshape(2, 3) merged = np.concatenate([a, b], axis=0) # 沿行方向摞高 print(merged) # [[ 0 1 2] # [ 3 4 5] # [ 6 7 8] # [ 9 10 11]] print(merged.shape) # (4, 3) print(np.shares_memory(merged, a)) # False —— 新内存,必然拷贝
代价意识:合并 2 份 100MB 的数组要新分配 200MB 并搬运,峰值瞬间三份并存。循环里反复合并是经典性能事故:
import numpy as np import time rng = np.random.default_rng(0) chunks = [rng.rand(1000, 100) for _ in range(200)] # 反面教材:循环 append 式合并,每次全量重拷 t0 = time.perf_counter() result = chunks[0] for c in chunks[1:]: result = np.concatenate([result, c]) # 第 k 次拷 k 份,总搬运量平方级 print("循环合并:", round(time.perf_counter() - t0, 3), "秒") # 约 0.05 秒 # 正确姿势:攒够一次性合并,只拷一遍 t0 = time.perf_counter() result2 = np.concatenate(chunks, axis=0) print("一次合并:", round(time.perf_counter() - t0, 4), "秒") # 约 0.001 秒 print(np.array_equal(result, result2)) # True
五十倍差距只是两百块的规模;块数上千时平方复杂度会让循环版彻底不可用。
stack 新增一个维度,把一组同形状数组"码成一摞":
import numpy as np x = np.array([1, 2, 3]) y = np.array([4, 5, 6]) print(np.stack([x, y], axis=0)) # [[1 2 3] # [4 5 6]] 形状 (2, 3) print(np.stack([x, y], axis=1)) # [[1 4] # [2 5] # [3 6]] 形状 (3, 2) print(np.concatenate([x, y])) # 拼接不加维 # [1 2 3 4 5 6] 形状 (6,)
hstack 与 vstack 是按"视觉方向"命名的快捷方式,但只在二维时直觉正确,一维时会反直觉:
import numpy as np p = np.array([1, 2]) q = np.array([3, 4]) print(np.vstack([p, q])) # (2,2):一维数组被当作行,先升维再拼 # [[1 2] # [3 4]] print(np.hstack([p, q])) # (4,):一维就是水平接龙 # [1 2 3 4] m = np.arange(4).reshape(2, 2) print(np.hstack([m, m]).shape) # (2, 4) 加宽 print(np.vstack([m, m]).shape) # (4, 2) 加高

import numpy as np data = np.arange(12) parts = np.split(data, 3) # 均分3份,不均会报错 print([p.tolist() for p in parts]) # [[0,1,2,3],[4,5,6,7],[8,9,10,11]] at = np.split(data, [4, 8]) # 按分割点位置切 print([p.tolist() for p in at]) # [[0,1,2,3],[4,5,6,7],[8,9,10,11]] uneven = np.array_split(data, 5) # 不均分也不报错 print([len(p) for p in uneven]) # [3, 3, 2, 2, 2] print(np.shares_memory(parts[0], data)) # True —— 切分是视图!
数据集切分是 split 的高频应用,完整模板:
import numpy as np rng = np.random.default_rng(6) X = rng.rand(1000, 20) # 特征 y = rng.integers(0, 2, 1000) # 标签 # 先洗牌再切分,两数组要同步洗:同一个排列作用于两者 perm = rng.permutation(len(X)) Xs, ys = X[perm], y[perm] # 花式索引拷贝,两份新内存但保持对齐 n_train = 700 X_train, X_rest = np.split(Xs, [n_train]) y_train, y_rest = np.split(ys, [n_train]) X_val, X_test = np.split(X_rest, [150]) y_val, y_test = np.split(y_rest, [150]) print(X_train.shape, X_val.shape, X_test.shape) # (700,20) (150,20) (150,20) print(np.bincount(y_train)) # 标签分布 [353 347] 大致均衡
洗牌用同一个 perm 作用两组数据,保证 X 与 y 行对行对齐——这是切分事故的第一大来源。注意 Xs 是花式索引的拷贝,随后的 split 才能给视图。
⚠️ 常见坑:np.split 均分不成立时抛 ValueError,很多人以为数据坏了;其实是 10 份切 3 份除不尽。要么用 array_split,要么用分割点列表明确指定边界。
下一节解决数据的进出:npy 与文本的速度鸿沟、memmap 如何处理比内存还大的文件。