Phase 03: Deep Learning Core

معرفی به JAX

پایتورچ تنسورها را جهش می دهد. تنسور فلو نمودارها را ایجاد می کند. جاکس عملکردهای خالص را جمع آوری می کند. آخرین یکی نحوه تفکر شما در مورد یادگیری عمیق را تغییر می دهد.

Type: Build

Languages: Python

Prerequisites: Phase 03 Lessons 01-10, basic NumPy

Time: ~90 minutes

اهداف یادگیری

  • کد شبکه عصبی عملکرد خالص را با استفاده از API عملکردی JAX بنویسید (jax.numpy، jax.grad، jax.jit، jax.vmap)
  • تفاوت اصلی طراحی بین جهش مشتاق PyTorch و مدل کامپیلیشن عملکردی JAX را توضیح دهید
  • استفاده از مجموعه jit و ویکتور سازی vmap برای سرعت بخشیدن به حلقه های آموزش در مقایسه با پایتون ساده
  • آموزش یک شبکه ساده در JAX و مقایسه مدیریت دولت صریح با رویکرد هدفمند PyTorch

مشکل

تو می دونی چطور شبکه های عصبی رو در PyTorch بسازیnn.Module، تماس بگیرید.backward()،و به سمت بهینه ساز فشار بده ، کار ميکنه . ميليون ها نفر ازش استفاده ميکنن

اما PyTorch یک محدودیت در DNA اش پخته شده است: آن را به دنبال عملیات مشتاقانه، یکی به یک زمان، در پایتون.tensor + tensorهر مرحله آموزش دوباره همان کد پایتون را تفسیر می کند. این کار خوب کار می کند تا زمانی که شما نیاز به آموزش یک مدل پارامتر 540 میلیارد در 2048 TPU دارید. سپس هزینه های بالای شما را می کشد.

گوگل DeepMind دوقلوها را بر روی JAX آموزش می دهد. آنترپیک کلود را بر روی JAX آموزش می دهد. این عملیات های کوچک نیستند - این بزرگترین عملیات آموزش شبکه عصبی در زمین هستند. آنها JAX را انتخاب کردند زیرا این چرخه آموزش شما را به عنوان یک برنامه قابل مرتب می کند، نه یک ردیف تماس های پایتون.

JAX با سه ابرقدرت: تفاوت خودکار، جمع آوری JIT به XLA و ویکتور سازی خودکار است. شما یک تابع را می نویسید که یک مثال را پردازش می کند. JAX به شما یک تابع را می دهد که یک دسته را پردازش می کند، گرادیانت ها را محاسبه می کند، به کد ماشین را مرتب می کند و در چندین دستگاه اجرا می کند. همه بدون تغییر عملکرد اصلی.

مفهوم

فلسفه جاکس

JAX يه چارچوبي فعاله بدون کلاس، بدون حالت متغير، نه.backward()روش.به جایش:

PyTorchJAX
nn.Module class with statePure function: f(params, x) -> y
loss.backward()jax.grad(loss_fn)(params, x, y)
Eager executionJIT compilation via XLA
for x in batch: manual loopjax.vmap(f) auto-vectorization
DataParallel / FSDPjax.pmap(f) auto-parallelism
Mutable model.parameters()Immutable pytree of arrays

این یک انتخاب سبک نیست. این یک محدودیت کامپایلر است. جمع آوری JIT نیازمند عملکردهای خالص است - ورودی های مشابه همیشه نتایج مشابهی را تولید می کنند، هیچ عوارض جانبی وجود ندارد. این محدودیت چیزی است که باعث می شود سرعت 100 برابر ممکن باشد.

jax.numpy: سطح آشنا

JAX API NumPy را در تسریع کننده ها مجدداً پیاده سازی می کند:

pythonimport jax.numpy as jnp

a = jnp.array([1.0, 2.0, 3.0])
b = jnp.array([4.0, 5.0, 6.0])
c = jnp.dot(a, b)

اسم هاي تابع همديگه، قوانين پخش همديگه، همديگه سيمانتیک برش دادن، اما آرایه ها روي GPU/TPU زنده هستند و هر عمل توسط کامپايلر قابل ردیابی است

يه تفاوت مهم: آرایه هاي JAX غير قابل تغيير هستنa[0] = 5. در عوض:a = a.at[0].set(5)این یک هفته عجیب و غریب است، سپس می کند -- تغییر ناپذیر بودن چیزی است که تحولات را شبیه بهgrad،jitوvmapقابل تدوین

jax.grad: خودکشی عملکردی

پیتورچ گرادینت ها را به تنسورها متصل می کند (.gradجاکس گرادینت ها را به تابع ها متصل می کند.

pythonimport jax

def f(x):
    return x ** 2

df = jax.grad(f)
df(3.0)

jax.gradیک تابع را می گیرد و یک تابع جدید را که گرادینت را محاسبه می کند، باز می آورد..backward()هیچ گراف محاسباتی در تنسورها ذخیره نشده است. گرادینت فقط یک تابع دیگر است که می توانید آن را فرا بگیرید، یا ترکیب کنید یا JIT-کمپیل کنید.

این به طور تعسفی تشکیل می شود:

pythond2f = jax.grad(jax.grad(f))
d2f(3.0)

مشتقات دوم، مشتقات سوم، جاکوبیان، هسیان، همه با ترکیبgrad.پایتورچ هم می تونه این کار رو بکنهtorch.autograd.functional.hessianدر جاکس، این پایه است.

محدودیت:gradهیچ گونه بیان نامه ای در داخل (آن ها در هنگام ردیابی اجرا می شوند، نه اجرا می شوند) هیچ جهش در حالت خارجی. هیچ تولید شماره تصادفی بدون مدیریت کلید صریح نیست.

jit: به XLA کامپایل کنید

python@jax.jit
def train_step(params, x, y):
    loss = loss_fn(params, x, y)
    return loss

fast_step = jax.jit(train_step)

در اولین تماس، JAX عملکرد را ردیابی می کند - آن را ثبت می کند که کدام عملیات انجام می شود، بدون اجرای آنها. سپس آن را به XLA (الجهبری خطی سریع) ، کمپایلر Google برای TPU ها و GPU ها می رساند. XLA عملیات را ترکیب می کند، کپی های حافظه اضافی را از بین می برد و کد ماشین بهینه سازی شده تولید می کند.

تماس های بعدی به طور کامل از پایتون خارج می شوند. کد مرتب شده در سرعت افزونه C ++ اجرا می شود.

وقتی JIT کمک می کند:

  • مراحل آموزش (همین محاسبه هزاران بار تکرار می شود)
  • انفرنس (مثل مدل، ورودی های مختلف)
  • هر تابع که بیش از یک بار با ورودی های مشابه نامیده شود

وقتی JIT درد میکنه:

  • عملکردهای با جریان کنترل پایتون که به ارزش ها بستگی دارد (if x > 0جایی که x یک ردیف ردیابی است)
  • محاسبه های یکبار (آموزش های مرتب شده بیش از زمان اجرا)
  • بازیابی (تراسیج اجرای واقعی را پنهان می کند)

محدودیت جریان کنترل واقعیهjax.lax.condجایگزین می شودif/else.jax.lax.scanجایگزین می شودforحلقه ها. این ها اختیاری نیستند. این ها قیمت جمع آوری هستند.

vmap: متریزیشن خودکار

شما یک تابع را می نویدید که یک مثال را پردازش می کند:

pythondef predict(params, x):
    return jnp.dot(params['w'], x) + params['b']

vmapآن را برای پردازش یک دسته بلند می کند:

pythonbatch_predict = jax.vmap(predict, in_axes=(None, 0))

in_axes=(None, 0)روش: دسته بندی نکنید params(مشترکه) ، دسته ای از محور 0 از x. هيچ راهنماييfor.چاپ .بدون تغییر شکل .بدون رشته ابعاد دسته بندی .جاکس ابعاد دسته را مشخص می کند و کل محاسبات را متری می کند

اين قند سنتکسي نيستvmapکد متریزه ای مخلوط تولید می کند که 10-100 برابر سریعتر از یک حلقه پایتون اجرا می شود. و این ترکیب با jitوgrad:

pythonper_example_grads = jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0))

هر نمونه يه خط اين تقريباً بدون هک ها در PyTorch ناممکن است

pmap: موازی داده ها در دستگاه ها

pythonparallel_step = jax.pmap(train_step, axis_name='devices')

pmapتکرار تابع در تمام دستگاه های موجود (GPU / TPU) و تقسیم دسته. در داخل تابع، jax.lax.pmeanوjax.lax.psumهم وقت سازی گرادینت ها در دستگاه ها

گوگل با استفاده از هزاران تراشه TPU v5e، دوقلوها را آموزش می دهدpmap(و جانشینش)shard_map) مدل برنامه نویسی: نسخه ی دستگاه یگانه را بنویسید، با pmap، تموم شد

پیترز: ساختار داده جهانی

جاکس روی "پایتری" کار می کند - ترکیبی از لیست ها، توپل ها، دیکت ها و آرایه ها. پارامترهای مدل شما یک پایتری هستند:

pythonparams = {
    'layer1': {'w': jnp.zeros((784, 256)), 'b': jnp.zeros(256)},
    'layer2': {'w': jnp.zeros((256, 128)), 'b': jnp.zeros(128)},
    'layer3': {'w': jnp.zeros((128, 10)),  'b': jnp.zeros(10)},
}

هر تحول جاکسgrad،jit،vmap-مي داند چطور از درختان پيتر عبور کنهjax.tree.map(f, tree)اعمال می شودfاین روش است که بهینه سازی کننده ها تمام پارامترها را به یک بار به روز می کنند:

pythonparams = jax.tree.map(lambda p, g: p - lr * g, params, grads)

نه.parameters()روش. بدون ثبت پارامتر. ساختار درخت مدل است.

عملکردی در مقابل هدفمند

فروشگاه های PyTorch در داخل اشیاء می گویند:

pythonclass Model(nn.Module):
    def __init__(self):
        self.linear = nn.Linear(784, 10)

    def forward(self, x):
        return self.linear(x)

JAX از تابع های خالص با حالت صریح استفاده می کند:

pythondef predict(params, x):
    return jnp.dot(x, params['w']) + params['b']

پارام ها منتقل می شوند. هیچ چیز ذخیره نمی شود. هیچ چیز جهش نمی یابد. این باعث می شود که هر عملکرد قابل آزمایش، قابل ترکیب و قابل تجمع باشد. همچنین به این معنی است که شما خود پارام ها را مدیریت می کنید - یا از یک کتابخانه مانند فلان یا یکوینوکس استفاده کنید.

اکوسیستم جی ای ایکس

جاکس به شما ابتدایی ها می دهد کتابخانه ها به شما ارگونومی می دهند:

LibraryRoleStyle
Flax (Google)Neural network layersnn.Module with explicit state
Equinox (Patrick Kidger)Neural network layersPytree-based, Pythonic
Optax (DeepMind)Optimizers + LR schedulesComposable gradient transforms
Orbax (Google)CheckpointingSave/restore pytrees
CLU (Google)Metrics + loggingTraining loop utilities

Optax کتابخانه بهینه سازی استاندارد است. این تغییر گرادینت (آدم، SGD، کپی) را از بروزرسانی پارامتر جدا می کند، بنابراین ترکیب آن معمولی است:

pythonoptimizer = optax.chain(
    optax.clip_by_global_norm(1.0),
    optax.adam(learning_rate=1e-3),
)

چه زمانی باید JAX vs PyTorch را استفاده کنید

FactorJAXPyTorch
TPU supportFirst-class (Google built both)Community-maintained (torch_xla)
GPU supportGood (CUDA via XLA)Best-in-class (native CUDA)
DebuggingHard (tracing + compilation)Easy (eager, line-by-line)
EcosystemResearch-focused (Flax, Equinox)Massive (HuggingFace, torchvision, etc.)
HiringNiche (Google/DeepMind/Anthropic)Mainstream (everywhere)
Large-scale trainingSuperior (XLA, pmap, mesh)Good (FSDP, DeepSpeed)
Prototyping speedSlower (functional overhead)Faster (mutate and go)
Production inferenceTensorFlow Serving, Vertex AITorchServe, Triton, ONNX
Who uses itDeepMind (Gemini), Anthropic (Claude)Meta (Llama), OpenAI (GPT), Stability AI

پاسخ صادقانه: از PyTorch استفاده کنید مگر اینکه دلیل خاصی برای استفاده از JAX داشته باشید. این دلایل این هستند: دسترسی به TPU، نیاز به هر مثال gradients، آموزش چند دستگاه در مقیاس گسترده، یا کار در گوگل / DeepMind / Anthropic.

شماره های تصادفی در JAX

JAX دارای حالت تصادفی جهانی نیست. هر عملیات تصادفی نیازمند کلید PRNG صریح است:

pythonkey = jax.random.PRNGKey(42)
key1, key2 = jax.random.split(key)
w = jax.random.normal(key1, shape=(784, 256))

این در ابتدا ناراحت کننده است. اما این قابلیت بازیافت در دستگاه ها و مجموعه ها را تضمین می کند. یک ویژگی که PyTorchtorch.manual_seedنمی تواند در تنظیمات چند GPU تضمین کند.

آن را بسازید

مرحله اول: تنظیم و داده ها

ما با استفاده از JAX و Optax یک MLP سه لایه را در MNIST آموزش می دهیم. 784 ورودی، دو لایه مخفی از 256 و 128 نورون، 10 کلاس خروجی.

pythonimport jax
import jax.numpy as jnp
from jax import random
import optax

def get_mnist_data():
    from sklearn.datasets import fetch_openml
    mnist = fetch_openml('mnist_784', version=1, as_frame=False, parser='auto')
    X = mnist.data.astype('float32') / 255.0
    y = mnist.target.astype('int')
    X_train, X_test = X[:60000], X[60000:]
    y_train, y_test = y[:60000], y[60000:]
    return X_train, y_train, X_test, y_test

مرحله دوم: شروع کردن پارامترها

هیچ کلاس نیست فقط یک تابع که یک پیتر را باز می گرداند:

pythondef init_params(key):
    k1, k2, k3 = random.split(key, 3)
    scale1 = jnp.sqrt(2.0 / 784)
    scale2 = jnp.sqrt(2.0 / 256)
    scale3 = jnp.sqrt(2.0 / 128)
    params = {
        'layer1': {
            'w': scale1 * random.normal(k1, (784, 256)),
            'b': jnp.zeros(256),
        },
        'layer2': {
            'w': scale2 * random.normal(k2, (256, 128)),
            'b': jnp.zeros(128),
        },
        'layer3': {
            'w': scale3 * random.normal(k3, (128, 10)),
            'b': jnp.zeros(10),
        },
    }
    return params

اون شروع به کار ميکنه، دستي انجام ميده، سه تا کليد PRNG از يك دانه جدا شده هر وزن يه تشکيل غير قابل تغيير در يک فرماني هست

مرحله سوم: عبور جلو

pythondef forward(params, x):
    x = jnp.dot(x, params['layer1']['w']) + params['layer1']['b']
    x = jax.nn.relu(x)
    x = jnp.dot(x, params['layer2']['w']) + params['layer2']['b']
    x = jax.nn.relu(x)
    x = jnp.dot(x, params['layer3']['w']) + params['layer3']['b']
    return x

def loss_fn(params, x, y):
    logits = forward(params, x)
    one_hot = jax.nn.one_hot(y, 10)
    return -jnp.mean(jnp.sum(jax.nn.log_softmax(logits) * one_hot, axis=-1))

عملکرد خالص، پارامز وارد، پیش بینی خارجself، هیچ حالت ذخیره نشدهloss_fnاین در واقع یک مقدار بسیار زیاد است.

مرحله چهارم: مرحله آموزش ی JIT

python@jax.jit
def train_step(params, opt_state, x, y):
    loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    return params, opt_state, loss

@jax.jit
def accuracy(params, x, y):
    logits = forward(params, x)
    preds = jnp.argmax(logits, axis=-1)
    return jnp.mean(preds == y)

jax.value_and_gradهر دو ارزش از دست دادن و گرادینت ها را در یک گذر باز می کند.@jax.jitبعد از اولین تماس، هر مرحله آموزش بدون لمس پایتون اجرا می شود.

مرحله پنجم: چرخه آموزش

pythonoptimizer = optax.adam(learning_rate=1e-3)

X_train, y_train, X_test, y_test = get_mnist_data()
X_train, X_test = jnp.array(X_train), jnp.array(X_test)
y_train, y_test = jnp.array(y_train), jnp.array(y_test)

key = random.PRNGKey(0)
params = init_params(key)
opt_state = optimizer.init(params)

batch_size = 128
n_epochs = 10

for epoch in range(n_epochs):
    key, subkey = random.split(key)
    perm = random.permutation(subkey, len(X_train))
    X_shuffled = X_train[perm]
    y_shuffled = y_train[perm]

    epoch_loss = 0.0
    n_batches = len(X_train) // batch_size
    for i in range(n_batches):
        start = i * batch_size
        xb = X_shuffled[start:start + batch_size]
        yb = y_shuffled[start:start + batch_size]
        params, opt_state, loss = train_step(params, opt_state, xb, yb)
        epoch_loss += loss

    train_acc = accuracy(params, X_train[:5000], y_train[:5000])
    test_acc = accuracy(params, X_test, y_test)
    print(f"Epoch {epoch + 1:2d} | Loss: {epoch_loss / n_batches:.4f} | "
          f"Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f}")

10 دوره. ~ 97٪ دقت آزمون. اولین دوره آهسته است (توسعه JIT). 2-10 دوره سریع است.

توجه کن که چه چیزی از دست رفته: نه.zero_grad()نه.backward()نه.step()تمام بروزرسانی یک تماس تابع ترکیب است. درجه بندی ها محاسبه می شوند، توسط آدم تبدیل می شوند و به پارامترها اعمال می شوند - همه در داخلtrain_step. .

ازش استفاده کن

فله: استاندارد گوگل

فلانس رایج ترین کتابخانه شبکه عصبی JAX است.nn.Moduleپس از آن، اما با مدیریت صریح دولت:

pythonimport flax.linen as nn

class MLP(nn.Module):
    @nn.compact
    def __call__(self, x):
        x = nn.Dense(256)(x)
        x = nn.relu(x)
        x = nn.Dense(128)(x)
        x = nn.relu(x)
        x = nn.Dense(10)(x)
        return x

model = MLP()
params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 784)))
logits = model.apply(params, x_batch)

ساختار مشابه با "پایتورچ" هست، اماparamsاز مدل جدا شده است. model.init()پارامز رو ایجاد ميکنهmodel.apply(params, x).مطابق جلو رو اجرا ميکنه .جسم مدل حالت نداره

شبابر: جایگزین پیتون

یک شبه عصر (از طرف پاتریک کیجر) مدل ها را به عنوان پایتری نشان می دهد:

pythonimport equinox as eqx

model = eqx.nn.MLP(
    in_size=784, out_size=10, width_size=256, depth=2,
    activation=jax.nn.relu, key=jax.random.PRNGKey(0)
)
logits = model(x)

خود مدل یک درخت است..apply()پارامترها فقط برگ های مدل هستند این به طرز فکر جیاکس نزدیک تر است

Optax: بهینه سازی سازنده

Optax تغییر گرادینت را از بروزرسانی جدا می کند:

pythonschedule = optax.warmup_cosine_decay_schedule(
    init_value=0.0, peak_value=1e-3,
    warmup_steps=1000, decay_steps=50000
)

optimizer = optax.chain(
    optax.clip_by_global_norm(1.0),
    optax.adamw(learning_rate=schedule, weight_decay=0.01),
)

کپی گریادیوتی، افزایش سرعت یادگیری، کاهش وزن همه اینها به عنوان یک زنجیره تحولات تشکیل شده است. هر تحول گرادیوتی را می بیند، آن ها را تغییر می دهد و آن ها را به بعدی منتقل می کند. هیچ کلاس بهینه سازی یکگانه ای وجود ندارد.

-باده

Installation:

bashpip install jax jaxlib optax flax

برای پشتیبانی GPU:

bashpip install jax[cuda12]

برای TPU (گووگلب کلاو):

bashpip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html

Performance gotchas:

  • اولین تماس JIT آهسته است (توسعه). قبل از مقایسه، گرم شوید.
  • از حلقه هاي پايتون روي آرایه هاي JAX داخل JIT اجتناب کنيد.jax.lax.scanیاjax.lax.fori_loop. .
  • jax.debug.print()کار در داخل JIT.print()نه، نه
  • پروفایل با jax.profilerیا TensorBoard. مجموعه XLA می تواند گلو های بطری را پنهان کند.
  • JAX 75 درصد حافظه GPU رو بطور پیش فرض اختصاص ميدهXLA_PYTHON_CLIENT_PREALLOCATE=falseبرای غیرفعال کردن

Checkpointing:

pythonimport orbax.checkpoint as ocp
checkpointer = ocp.PyTreeCheckpointer()
checkpointer.save('/tmp/model', params)
restored = checkpointer.restore('/tmp/model')

This lesson produces:

  • outputs/prompt-jax-optimizer.md-- یک دستور برای انتخاب تنظیمات بهینه سازی JAX مناسب
  • outputs/skill-jax-patterns.md-- يه مهارتي که الگوهای فعالي را در JAX پوشش ميده

تمرینات

  1. در JAX، ترک کردن نیاز به کلید PRNG دارد - یک کلید را از طریق گذرگاه جلو به هم بزنید و آن را برای هر لایه ترک کردن تقسیم کنید. دقت آزمون را با و بدون مقایسه کنید.
  1. استفاده کنیدjax.vmapبرای محاسبه هر نمونه گرادینت برای یک دسته از 32 تصویر MNIST. برای هر مثال نرمال گرادینت را محاسبه کنید. کدام نمونه ها بیشترین گرادینت را دارند و چرا؟
  1. عملکرد پیش رو دستی را با یک عمومی جایگزین کنید mlp_forward(params, x)که برای هر تعداد لایه ای کار می کند.jax.tree.leavesبرای تعیین عمق به طور خودکار.
  1. مرحله آموزش را با و بدون آن بنچ مارک کنید @jax.jit. زمان 100 مرحله هر کدام چقدر سرعت افزايش سخت افزاري شما بزرگ است؟
  1. از طریق ترکیب کردن کتیج گرادینت را اجرا کنید optax.chain(optax.clip_by_global_norm(1.0), optax.adam(1e-3)). تمرین با و بدون برش. نماد گرادینت رو بر روی تمرین نشان بده تا تا اثر رو ببین

اصطلاحات کلیدی

TermWhat people sayWhat it actually means
XLA"The thing that makes JAX fast"Accelerated Linear Algebra -- a compiler that fuses operations and generates optimized GPU/TPU kernels from a computation graph
JIT"Just-in-time compilation"JAX traces the function on first call, compiles to XLA, then runs the compiled version on subsequent calls
Pure function"No side effects"A function where the output depends only on inputs -- no global state, no mutation, no randomness without explicit keys
vmap"Auto-batching"Transforms a function that processes one example into one that processes a batch, without rewriting
pmap"Auto-parallelism"Replicates a function across multiple devices and splits the input batch
Pytree"Nested dict of arrays"Any nested structure of lists, tuples, dicts, and arrays that JAX can traverse and transform
Tracing"Recording the computation"JAX executes the function with abstract values to build a computation graph, without computing real results
Functional autodiff"grad of a function"Computing derivatives by transforming functions, not by attaching gradient storage to tensors
Optax"JAX's optimizer library"A composable library of gradient transformations -- Adam, SGD, clipping, scheduling -- that chain together
Flax"JAX's nn.Module"Google's neural network library for JAX, adding layer abstractions while keeping state explicit

خواندن بیشتر

  • اسناد JAX: https://jax.readthedocs.io/-- دکترای رسمی، با آموزش های عالی در مورد Graduate، jit و vmap
  • "JAX: تحولات قابل ترکیب برنامه های پایتون+نمپی" (برادبری و همکارانش، 2018) - مقاله اصلی که فلسفه طراحی را توضیح می دهد
  • اسناد فلان: https://flax.readthedocs.io/-- کتابخانه شبکه عصبی گوگل برای JAX
  • پاتریک کیجر، "اقیاس: شبکه های عصبی در JAX از طریق PyTrees قابل تماس و تحولات فیلتر شده" (2021) -- جایگزین پایتونیک برای فله
  • DeepMind، "Optax: تبدیل و بهینه سازی گرادینت ترکیب شده" -- کتابخانه بهینه سازی استاندارد
  • "تو نمی دانی که جاکس" (کولین رافل، 2020) - راهنمای عملی برای جاکس گات ها و الگوهای، از یکی از نویسندگان T5

This free lesson is part of the AI Engineering from Scratch curriculum. Read the full explanation, run the lesson code, and verify the result in the interactive reader or from the repository source.

Browse the complete course catalog or open this lesson on GitHub.