#40416·jax

XLA 变化显著地增加了计算时间

作者: JanLuca创建于 2026年9月3日更新于 2026年9月16日
标签bug

说明

openxla/xla@d97f3a7bffe359846e45948e2cab8a6c9be2993a(投入54ba3e02012afe6663943455ce57fea0f8d2cddd)的XLA变化,增加了以下例子的编译时间(使用基于JAX的 variPEPS/ variPEPS Python@465915c22ed1a17f14c17c17c6c6e28ca4985e13bd7库),增加了~5.

在承诺54ba3e02012afe66639435ce57fea0f8d2cdd时,仍有可能减少使用QQLA FLAGS='-xla cpu eximental enable tiling propagation=false'的编译时间的增加. 由于XLA在openxla/xla@8d14527c7fddf08f994d48bd696a84db26064394(投入JAX,承诺4c7da2781a48620f33386f1084ab9069b80f03b9. 更改后xla旗不再影响编译时间.

示例

导入时间

导入变量
导入 Jax 。 数字为jnp

Id = jnp.eye(2个)
Sx = jnp.array ([[0, 1], [1, 0]]) / 2.
Sy = jnp.array ([0, -1j], [1j, 0]]) / 2.
Sz = jnp.array ([[1, ], [0,-1]) / 2 (中文(简体) ).

闸口= (jnp.kron (Sx, Sx) + jnp.kron (Sy, Sy) + jnp.kron (Sz, Sz)).

varipeps.config.config.ctmrg print steps=真实
varipeps.config.config.ctmrg full projector method=varipeps.config.Projector Method.fulL QR
varipeps.config.ctmrg convergence eps=1 * 1e-7
varipeps.config.ad custom print steps=为真名.
varipeps.config.ad custom fixed point 方法=varipeps.config. Grad Fixed Point 方法. ITERATIONT.
varpeps.config.ctmrg heuristic development chi = 虚假
varipeps.config.ctmrg 增加 truncation eps=虚假
varipeps.config.ctmrg heuristic 增加 chi=虚假

单位细胞 = varipeps.peps.PEPS Unit Cell.random()
((0, 1, 1, 0)),
2,二.
2,二.
20岁,
浮点
20岁,
varipeps.peps.PEPS Type.SQUARE, (中文(简体) ).
种子数=6396316583739234,
(中文(简体) ).

# 测量执行时间,包括叮当
开始=时间.perf counter ()
结果,   = varipeps.ctmrg.calc ctmrg env(tuple(i.tensor for i in unitcell.get unique tensors ())), unitcell, 强制执行 elementwise convergence=False)
打印( time.perf  counter () - start)

QQ系统信息(Python版本,Jaxlib版本,加速器等).

jax:0.11.dev20260726+4c7da2781a (中文(简体) ).
jaxlib: 0.11.dev0+ 自建
数字:2.5.2
Python: 3.14.5 (主,2026年5月10日,19:28:16) [Clang 22.1.3] (中文(简体) ).
设备信息: cpu-1, 1个本地设备"
进程( C): 1
平台:uname result(系统='Linux',节点='qmio20',发布='6.12.101+deb13-amd64',版本=' 1 SMP PREEMPT DYNAMIC Debian 6.12.101-1',机器='x86 64')