已合并
fix: uniform 大 tensor 分块生成降内存(复用 CHUNK_ELEMS=4M,与单次全量逐位相同) #213
dengguojie创建于 28 天前
fix: uniform 大 tensor 分块生成降内存(复用 CHUNK_ELEMS=4M,与单次全量逐位相同) #213
已合并
共 2 个文件变更+130-2
| @@ -181,3 +181,113 @@ def test_normal_chunked_generate_dtype_preserved(): | |||
| 181 | n = 2 * CHUNK_ELEMS + 8 | 181 | n = 2 * CHUNK_ELEMS + 8 |
| 182 | arr = RandomData("bfloat16", (n,), (-1.0, 1.0)).generate(distribution="normal") | 182 | arr = RandomData("bfloat16", (n,), (-1.0, 1.0)).generate(distribution="normal") |
| 183 | assert str(arr.dtype) == "bfloat16" | 183 | assert str(arr.dtype) == "bfloat16" |
| 184 | + | ||
| 185 | + | ||
| 186 | + | ||
| 187 | + "dtype", | ||
| 188 | + ["float16", "bfloat16", "float32", "float64", "int32", "int64", "uint8", "bool"], | ||
| 189 | +) | ||
| 190 | +def test_uniform_chunked_bitwise_equal(dtype): | ||
| 191 | + """uniform 大 tensor 分块生成与单次全量生成逐位相同(含非整除尾块)。""" | ||
| 192 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 193 | + from ttk.utilities.dtypes import resolve_custom_numpy_dtypes | ||
| 194 | + | ||
| 195 | + n = 2 * CHUNK_ELEMS + 1024 | ||
| 196 | + resolved = resolve_custom_numpy_dtypes([dtype])[0] | ||
| 197 | + | ||
| 198 | + np.random.seed(42) | ||
| 199 | + full = np.random.uniform(-1.0, 1.0, n).astype(resolved, copy=False) | ||
| 200 | + np.random.seed(42) | ||
| 201 | + chunked = RandomData._gen_uniform_data(-1.0, 1.0, resolved, (n,)) | ||
| 202 | + assert chunked.dtype == full.dtype | ||
| 203 | + assert np.array_equal(chunked, full) | ||
| 204 | + | ||
| 205 | + | ||
| 206 | +def test_uniform_small_tensor_keeps_single_shot_path(): | ||
| 207 | + """elem_count <= 2*CHUNK_ELEMS 时保持单次全量路径(行为与改动前一致)。""" | ||
| 208 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 209 | + | ||
| 210 | + n = 2 * CHUNK_ELEMS | ||
| 211 | + | ||
| 212 | + np.random.seed(42) | ||
| 213 | + full = np.random.uniform(-1.0, 1.0, n).astype("float32", copy=False) | ||
| 214 | + np.random.seed(42) | ||
| 215 | + small = RandomData._gen_uniform_data(-1.0, 1.0, "float32", (n,)) | ||
| 216 | + assert np.array_equal(small, full) | ||
| 217 | + | ||
| 218 | + | ||
| 219 | +def test_uniform_chunked_generate_dtype_preserved(): | ||
| 220 | + """generate 入口走 uniform 分块路径后 dtype 保持声明值(防 float64 提升/降级)。""" | ||
| 221 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 222 | + | ||
| 223 | + n = 2 * CHUNK_ELEMS + 8 | ||
| 224 | + arr = RandomData("float32", (n,), (-1.0, 1.0)).generate() | ||
| 225 | + assert str(arr.dtype) == "float32" | ||
| 226 | + | ||
| 227 | + | ||
| 228 | + | ||
| 229 | + "tail", | ||
| 230 | + [ | ||
| 231 | + 1, | ||
| 232 | + 4_000_000 - 1, | ||
| 233 | + ], | ||
| 234 | +) | ||
| 235 | +def test_uniform_chunked_tail_boundaries_bitwise_equal(tail): | ||
| 236 | + """非 4M 对齐边界:最小分块入口(2xCHUNK+1)与最大尾块(CHUNK-1)均逐位相同。""" | ||
| 237 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 238 | + | ||
| 239 | + n = 2 * CHUNK_ELEMS + tail | ||
| 240 | + | ||
| 241 | + np.random.seed(42) | ||
| 242 | + full = np.random.uniform(-1.0, 1.0, n).astype("float32", copy=False) | ||
| 243 | + np.random.seed(42) | ||
| 244 | + chunked = RandomData._gen_uniform_data(-1.0, 1.0, "float32", (n,)) | ||
| 245 | + assert np.array_equal(chunked, full) | ||
| 246 | + | ||
| 247 | + | ||
| 248 | +def test_uniform_chunked_exact_multiple_bitwise_equal(): | ||
| 249 | + """整除边界(n = 3*CHUNK,无尾块)分块生成与单次全量逐位相同。""" | ||
| 250 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 251 | + | ||
| 252 | + n = 3 * CHUNK_ELEMS | ||
| 253 | + | ||
| 254 | + np.random.seed(42) | ||
| 255 | + full = np.random.uniform(-1.0, 1.0, n).astype("float32", copy=False) | ||
| 256 | + np.random.seed(42) | ||
| 257 | + chunked = RandomData._gen_uniform_data(-1.0, 1.0, "float32", (n,)) | ||
| 258 | + assert np.array_equal(chunked, full) | ||
| 259 | + | ||
| 260 | + | ||
| 261 | +def test_uniform_chunked_actually_splits(monkeypatch): | ||
| 262 | + """大 tensor 必须真实走分块路径:uniform 调用次数 >= 2(防阈值被误改后静默退化为单发)。""" | ||
| 263 | + import numpy | ||
| 264 | + | ||
| 265 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 266 | + | ||
| 267 | + calls = [] | ||
| 268 | + real_uniform = numpy.random.uniform | ||
| 269 | + | ||
| 270 | + def counting_uniform(low, high, size): | ||
| 271 | + calls.append(size) | ||
| 272 | + return real_uniform(low, high, size) | ||
| 273 | + | ||
| 274 | + monkeypatch.setattr("ttk.utilities.data.numpy.random.uniform", counting_uniform) | ||
| 275 | + arr = RandomData("float32", (2 * CHUNK_ELEMS + 1,), (-1.0, 1.0)).generate() | ||
| 276 | + assert len(calls) >= 2 | ||
| 277 | + assert arr.dtype == numpy.dtype("float32") # dtype 保持 | ||
| 278 | + | ||
| 279 | + | ||
| 280 | +def test_uniform_chunked_multidim_bitwise_equal(): | ||
| 281 | + """多维 shape(总元素数非 4M 对齐)分块生成与单次全量逐位相同(ravel 写入顺序一致)。""" | ||
| 282 | + from ttk.utilities.data import CHUNK_ELEMS | ||
| 283 | + | ||
| 284 | + shape = (3, CHUNK_ELEMS // 2 + 17, 5) # 总数 = 1.5*CHUNK 非对齐,含多维 stride | ||
| 285 | + n = int(np.prod(shape)) | ||
| 286 | + assert n > 2 * CHUNK_ELEMS | ||
| 287 | + | ||
| 288 | + np.random.seed(42) | ||
| 289 | + full = np.random.uniform(-1.0, 1.0, shape).astype("float32", copy=False) | ||
| 290 | + np.random.seed(42) | ||
| 291 | + chunked = RandomData._gen_uniform_data(-1.0, 1.0, "float32", shape) | ||
| 292 | + assert chunked.shape == shape | ||
| 293 | + assert np.array_equal(chunked, full) | ||
| @@ -78,7 +78,7 @@ class RandomData: | |||
| 78 | low = 2**-3 | 78 | low = 2**-3 |
| 79 | if not numpy.isfinite(high) or high <= low: | 79 | if not numpy.isfinite(high) or high <= low: |
| 80 | high = 2**7 | 80 | high = 2**7 |
| 81 | - f32 = numpy.random.uniform(low, high, self._shape).astype("float32") | 81 | + f32 = self._gen_uniform_data(low, high, "float32", self._shape) |
| 82 | from .dtypes import numpy_float8_e8m0 | 82 | from .dtypes import numpy_float8_e8m0 |
| 83 | 83 | ||
| 84 | np_array = f32.astype(numpy_float8_e8m0()) | 84 | np_array = f32.astype(numpy_float8_e8m0()) |
| @@ -180,6 +180,24 @@ class RandomData: | |||
| 180 | flat[start:end] = gen.rvs(end - start) | 180 | flat[start:end] = gen.rvs(end - start) |
| 181 | return out | 181 | return out |
| 182 | 182 | ||
| 183 | + | ||
| 184 | + def _gen_uniform_data(low, high, dtype, shape): | ||
| 185 | + """Generate uniform samples into a pre-allocated typed buffer chunk by chunk. | ||
| 186 | + | ||
| 187 | + numpy.random.uniform consumes the RandomState stream element by element | ||
| 188 | + (low + (high-low) * next_double), so the chunked output is bitwise identical | ||
| 189 | + to a single full-size uniform(...).astype(dtype) call. | ||
| 190 | + """ | ||
| 191 | + elem_count = int(numpy.prod(shape)) if numpy.ndim(shape) else 1 | ||
| 192 | + if elem_count <= 2 * CHUNK_ELEMS: | ||
| 193 | + return numpy.random.uniform(low, high, shape).astype(dtype, copy=False) | ||
| 194 | + out = numpy.empty(shape, dtype=dtype) | ||
| 195 | + flat = out.ravel() | ||
| 196 | + for start in range(0, elem_count, CHUNK_ELEMS): | ||
| 197 | + end = min(start + CHUNK_ELEMS, elem_count) | ||
| 198 | + flat[start:end] = numpy.random.uniform(low, high, end - start) | ||
| 199 | + return out | ||
| 200 | + | ||
| 183 | def _random( | 201 | def _random( |
| 184 | self, dtype, shape: Union[list, tuple], is_complex_imag: bool = False, distribution: str = "uniform" | 202 | self, dtype, shape: Union[list, tuple], is_complex_imag: bool = False, distribution: str = "uniform" |
| 185 | ) -> numpy.ndarray: | 203 | ) -> numpy.ndarray: |
| @@ -228,7 +246,7 @@ class RandomData: | |||
| 228 | 246 | ||
| 229 | array = t.view(torch.uint16).numpy().view(np_bf16).reshape(shape) | 247 | array = t.view(torch.uint16).numpy().view(np_bf16).reshape(shape) |
| 230 | else: | 248 | else: |
| 231 | - array = numpy.random.uniform(low, high, shape).astype(dtype, copy=False) | 249 | + array = self._gen_uniform_data(low, high, dtype, shape) |
| 232 | return self._mix_expect_data(array, dtype, shape, is_complex_imag) | 250 | return self._mix_expect_data(array, dtype, shape, is_complex_imag) |
| 233 | 251 | ||
| 234 | def _mix_expect_data( | 252 | def _mix_expect_data( |