这个答案给出了一个使用structured arrays 的解决方案。它有以下要求: Ggven 一个函数 f 返回 N 数组,并且每个返回数组的大小可以不同 - 那么对于 f 的所有结果,len(array_i) 必须始终一样。例如。
arrs_a = f("a")
arrs_b = f("b")
for sub_arr_a, sub_arr_b in zip(arrs_a, arrs_b):
assert len(sub_arr_a) == len(sub_arr_b)
如果上述情况属实,那么您可以使用结构化数组。结构化数组就像普通数组一样,只是具有复杂的数据类型。例如,我可以指定一个数据类型,它由一个形状为5 的整数数组和另一个形状为(2, 2) 的浮点数数组组成。例如。
# define what a record looks like
dtype = [
# tuples of (field_name, data_type)
("a", "5i4"), # array of five 4-byte ints
("b", "(2,2)f8"), # 2x2 array of 8-byte floats
]
使用dtype,您可以创建一个结构化数组,并将所有结果一次性设置在结构化数组上。
import numpy as np
def func(n):
"mock implementation of func"
return (
np.ones(5) * n,
np.ones((2,2))* n
)
# define what a record looks like
dtype = [
# tuples of (field_name, data_type)
("a", "5i4"), # array of five 4-byte ints
("b", "(2,2)f8"), # 2x2 array of 8-byte floats
]
size = 5
# create array
arr = np.empty(size, dtype=dtype)
# fill in values
for i in range(size):
# func must return a tuple
# or you must convert the returned value to a tuple
arr[i] = func(i)
# alternate way of instantiating arr
arr = np.fromiter((func(i) for i in range(size)), dtype=dtype, count=size)
# How to use structured arrays
# access individual record
print(arr[1]) # prints ([1, 1, 1, 1, 1], [[1, 1], [1, 1]])
# access specific value -- get second record -> get b field -> get value at 0,0
assert arr[2]['b'][0,0] == 2
# access all values of a specific field
print(arr['a']) # prints all the a arrays