已合并
fix: uniform 大 tensor 分块生成降内存(复用 CHUNK_ELEMS=4M,与单次全量逐位相同) #213
dengguojie创建于 28 天前
fix: uniform 大 tensor 分块生成降内存(复用 CHUNK_ELEMS=4M,与单次全量逐位相同) #213
已合并
dengguojie创建于 28 天前
共 2 个文件变更+130-2
@@ -181,3 +181,113 @@ def test_normal_chunked_generate_dtype_preserved():
181 n = 2 * CHUNK_ELEMS + 8181 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+@pytest.mark.parametrize(
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+@pytest.mark.parametrize(
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**-378 low = 2**-3
79 if not numpy.isfinite(high) or high <= low:79 if not numpy.isfinite(high) or high <= low:
80 high = 2**780 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_e8m082 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 out181 return out
182 182 
183+ @staticmethod
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(