You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
# Pin the input to the CPU, which keeps the whole computation there
465
+
x0_cpu = jax.device_put(0.1, cpu)
466
+
467
+
@partial(jax.jit, static_argnames=("n",))
465
468
def qm_jax_fori(x0, n, α=4.0):
466
469
467
470
x = jnp.empty(n + 1).at[0].set(x0)
@@ -475,7 +478,7 @@ def qm_jax_fori(x0, n, α=4.0):
475
478
```
476
479
477
480
* ما `n` را ایستا نگه میداریم زیرا بر اندازه آرایه تأثیر میگذارد و از این رو JAX میخواهد روی مقدار آن در کد کامپایل شده تخصصی شود.
478
-
* ما به CPU از طریق `device=cpu`متصل میمانیم زیرا این بار کاری ترتیبی از بسیاری عملیات کوچک تشکیل شده است که فرصت کمی برای موازیسازی GPU باقی میگذارد.
481
+
* ما ورودی را با `jax.device_put` به CPU متصل میکنیم (که کل محاسبات را روی CPU نگه میدارد) زیرا این بار کاری ترتیبی از بسیاری عملیات کوچک تشکیل شده است که فرصت کمی برای موازیسازی GPU باقی میگذارد.
479
482
480
483
مهم: اگرچه `at[t].set` در هر مرحله ظاهراً یک آرایه جدید ایجاد میکند، در داخل یک تابع کامپایلشده با JIT، کامپایلر تشخیص میدهد که آرایه قدیمی دیگر مورد نیاز نیست و بهروزرسانی را در جا انجام میدهد!
481
484
@@ -484,7 +487,7 @@ def qm_jax_fori(x0, n, α=4.0):
484
487
```{code-cell} ipython3
485
488
with qe.Timer():
486
489
# First run
487
-
x_jax = qm_jax_fori(0.1, n)
490
+
x_jax = qm_jax_fori(x0_cpu, n)
488
491
# Hold interpreter
489
492
x_jax.block_until_ready()
490
493
```
@@ -494,7 +497,7 @@ with qe.Timer():
494
497
```{code-cell} ipython3
495
498
with qe.Timer():
496
499
# Second run
497
-
x_jax = qm_jax_fori(0.1, n)
500
+
x_jax = qm_jax_fori(x0_cpu, n)
498
501
# Hold interpreter
499
502
x_jax.block_until_ready()
500
503
```
@@ -508,7 +511,7 @@ JAX نیز برای این عملیات ترتیبی کاملاً کارآمد
508
511
این روش جایگزین، به طور قابل بحث، بیشتر با رویکرد تابعی JAX همسو است --- اگرچه سینتکس آن به خاطر سپردن دشواری دارد.
0 commit comments