Skip to content

Commit 6167963

Browse files
authored
🌐 [translation-sync] [numpy_vs_numba_vs_jax] Fix deprecated device= argument on jax.jit (#128)
* Update translation: lectures/numpy_vs_numba_vs_jax.md * Update translation: .translate/state/numpy_vs_numba_vs_jax.md.yml
1 parent 219f346 commit 6167963

2 files changed

Lines changed: 12 additions & 9 deletions

File tree

.translate/state/numpy_vs_numba_vs_jax.md.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
1-
source-sha: d08a73d48a409509d7d6f6585b99c2c8909c9a28
2-
synced-at: "2026-05-14"
1+
source-sha: d37b1d8adbf6e18b17e125cca761a6eb2ccd9041
2+
synced-at: "2026-06-19"
33
model: claude-sonnet-4-6
44
mode: UPDATE
55
section-count: 3

lectures/numpy_vs_numba_vs_jax.md

Lines changed: 10 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -461,7 +461,10 @@ Numba این عملیات ترتیبی را به طور بسیار کارآمد
461461
```{code-cell} ipython3
462462
cpu = jax.devices("cpu")[0]
463463
464-
@partial(jax.jit, static_argnames=("n",), device=cpu)
464+
# 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",))
465468
def qm_jax_fori(x0, n, α=4.0):
466469
467470
x = jnp.empty(n + 1).at[0].set(x0)
@@ -475,7 +478,7 @@ def qm_jax_fori(x0, n, α=4.0):
475478
```
476479

477480
* ما `n` را ایستا نگه می‌داریم زیرا بر اندازه آرایه تأثیر می‌گذارد و از این رو JAX می‌خواهد روی مقدار آن در کد کامپایل شده تخصصی شود.
478-
* ما به CPU از طریق `device=cpu` متصل می‌مانیم زیرا این بار کاری ترتیبی از بسیاری عملیات کوچک تشکیل شده است که فرصت کمی برای موازی‌سازی GPU باقی می‌گذارد.
481+
* ما ورودی را با `jax.device_put` به CPU متصل می‌کنیم (که کل محاسبات را روی CPU نگه می‌دارد) زیرا این بار کاری ترتیبی از بسیاری عملیات کوچک تشکیل شده است که فرصت کمی برای موازی‌سازی GPU باقی می‌گذارد.
479482

480483
مهم: اگرچه `at[t].set` در هر مرحله ظاهراً یک آرایه جدید ایجاد می‌کند، در داخل یک تابع کامپایل‌شده با JIT، کامپایلر تشخیص می‌دهد که آرایه قدیمی دیگر مورد نیاز نیست و به‌روزرسانی را در جا انجام می‌دهد!
481484

@@ -484,7 +487,7 @@ def qm_jax_fori(x0, n, α=4.0):
484487
```{code-cell} ipython3
485488
with qe.Timer():
486489
# First run
487-
x_jax = qm_jax_fori(0.1, n)
490+
x_jax = qm_jax_fori(x0_cpu, n)
488491
# Hold interpreter
489492
x_jax.block_until_ready()
490493
```
@@ -494,7 +497,7 @@ with qe.Timer():
494497
```{code-cell} ipython3
495498
with qe.Timer():
496499
# Second run
497-
x_jax = qm_jax_fori(0.1, n)
500+
x_jax = qm_jax_fori(x0_cpu, n)
498501
# Hold interpreter
499502
x_jax.block_until_ready()
500503
```
@@ -508,7 +511,7 @@ JAX نیز برای این عملیات ترتیبی کاملاً کارآمد
508511
این روش جایگزین، به طور قابل بحث، بیشتر با رویکرد تابعی JAX همسو است --- اگرچه سینتکس آن به خاطر سپردن دشواری دارد.
509512

510513
```{code-cell} ipython3
511-
@partial(jax.jit, static_argnames=("n",), device=cpu)
514+
@partial(jax.jit, static_argnames=("n",))
512515
def qm_jax_scan(x0, n, α=4.0):
513516
def update(x, t):
514517
x_new = α * x * (1 - x)
@@ -525,7 +528,7 @@ def qm_jax_scan(x0, n, α=4.0):
525528
```{code-cell} ipython3
526529
with qe.Timer():
527530
# First run
528-
x_jax = qm_jax_scan(0.1, n)
531+
x_jax = qm_jax_scan(x0_cpu, n)
529532
# Hold interpreter
530533
x_jax.block_until_ready()
531534
```
@@ -535,7 +538,7 @@ with qe.Timer():
535538
```{code-cell} ipython3
536539
with qe.Timer():
537540
# Second run
538-
x_jax = qm_jax_scan(0.1, n)
541+
x_jax = qm_jax_scan(x0_cpu, n)
539542
# Hold interpreter
540543
x_jax.block_until_ready()
541544
```

0 commit comments

Comments
 (0)