•10 min read

Học sâu với JAX

Học sâu với JAX

Trong bối cảnh các framework học máy đang phát triển nhanh chóng, JAX đã nổi lên như một công cụ mang tính cách mạng, định hình lại căn bản cách chúng ta xây dựng và mở rộng các phép tính số học. Được phát triển bởi Google Research, JAX không chỉ đơn thuần là một framework học sâu khác giống như TensorFlow hay PyTorch; thay vào đó, nó là một hệ thống mở rộng cho các phép biến đổi hàm có thể kết hợp, được xây dựng trên nền tảng NumPy tăng tốc. Bằng cách đưa tính năng tự động vi phân và biên dịch đúng lúc (JIT) trực tiếp vào các mảng NumPy, JAX mang đến một phương pháp học sâu sạch sẽ, thuần túy về mặt toán học, hoàn toàn phù hợp với các mô hình lập trình hàm. Trong bài phân tích kỹ thuật chuyên sâu toàn diện này, chúng ta sẽ khám phá các nguyên tắc cốt lõi làm cho JAX mạnh mẽ, xem xét hệ sinh thái của nó để phát triển mạng nơ-ron, và xây dựng sự hiểu biết về cách tạo ra các vòng lặp huấn luyện học sâu hiệu suất cao từ đầu.

Audio Briefing
0:00 / 0:00

Các Nguyên tắc Cốt lõi của JAX: Biến đổi Hàm

Cốt lõi của JAX là triết lý thiết kế tập trung vào các hàm thuần túy và các phép biến đổi có thể kết hợp. Không giống như các framework hướng đối tượng truyền thống, nơi các mô hình duy trì trạng thái nội bộ qua các bước huấn luyện, JAX khuyến khích thực thi không trạng thái. Các tính năng cốt lõi của JAX có thể truy cập được thông qua một vài phép biến đổi hàm cực kỳ mạnh mẽ có thể áp dụng cho các hàm Python tiêu chuẩn.

Biên dịch đúng lúc (jax.jit)

Trụ cột đầu tiên của JAX là khả năng biên dịch đúng lúc mã Python bằng XLA (Accelerated Linear Algebra). Khi bạn trang trí một hàm Python bằng @jax.jit, JAX sẽ theo dõi các hoạt động được thực hiện trên các mảng đầu vào và biên dịch chúng thành một chuỗi hoạt động XLA được tối ưu hóa. Quá trình này loại bỏ chi phí của trình thông dịch Python trong quá trình thực thi, cho phép mã chạy liền mạch và hiệu quả trên các bộ tăng tốc như GPU và TPU.

Điều làm cho jax.jit mạnh mẽ một cách độc đáo là khả năng kết hợp của nó. Bạn có thể dễ dàng biên dịch JIT một hàm mà bên trong gọi các hàm đã được biên dịch JIT khác, hoặc biên dịch toàn bộ bước huấn luyện của một mạng nơ-ron phức tạp thành một kernel nguyên khối, được tối ưu hóa duy nhất. Chiến lược biên dịch mạnh mẽ này thường mang lại những cải thiện đáng kể về hiệu suất so với các framework thực thi tức thời, mệnh lệnh.

Tự động vi phân (jax.grad)

Trụ cột thứ hai là tự động vi phân. Phép biến đổi jax.grad nhận một hàm có giá trị vô hướng và trả về một hàm mới tính toán đạo hàm của nó đối với các đối số của nó. Bởi vì JAX hoạt động trên các hàm thuần túy, việc tính toán đạo hàm mang lại cảm giác tự nhiên về mặt toán học. JAX hỗ trợ cả tự động vi phân chế độ tiến và chế độ ngược, và các phép biến đổi này có thể được kết hợp vô hạn. Bạn có thể tính toán đạo hàm bậc cao hơn chỉ bằng cách xâu chuỗi các lệnh gọi jax.grad (ví dụ: jax.grad(jax.grad(f))), điều này cực kỳ hữu ích cho các kỹ thuật tối ưu hóa nâng cao, học siêu cấp và mạng nơ-ron được thông tin vật lý.

Vector hóa (jax.vmap)

Trụ cột thứ ba là vector hóa tự động thông qua jax.vmap. Trong học sâu, chúng ta liên tục xử lý các lô dữ liệu. Theo truyền thống, điều này đòi hỏi phải viết lại mã một cách cẩn thận để hoạt động trên các tensor có chiều cao hơn. Với jax.vmap, bạn có thể viết một hàm hoạt động trên một điểm dữ liệu duy nhất và biến đổi nó thành một hàm tự động và hiệu quả hoạt động trên một lô. JAX đẩy vector hóa xuống cấp độ XLA, đảm bảo rằng phần cứng được sử dụng tối ưu mà không cần đến gánh nặng nhận thức của việc quản lý kích thước lô thủ công.

Song song hóa (jax.pmap và jax.sharding)

Để huấn luyện phân tán trên nhiều thiết bị (như nhiều GPU hoặc các pod TPU), JAX cung cấp các công cụ cho tính song song Single-Program Multiple-Data (SPMD). Phép biến đổi jax.pmap cho phép bạn nhân bản một hàm và thực thi nó đồng thời trên các thiết bị có sẵn, với các nguyên thủy giao tiếp tập thể (như jax.lax.pmean) được tích hợp sẵn để đồng bộ hóa gradient. Gần đây hơn, JAX đã giới thiệu các API phân chia mảng nâng cao cho phép kiểm soát chi tiết cách dữ liệu và tham số mô hình được phân phối trên các lưới thiết bị, cho phép mở rộng dễ dàng sang các mô hình lớn.

Advertisement

Xây dựng Mạng Nơ-ron: Hệ sinh thái Flax

Trong khi JAX cung cấp nền tảng số học, nó không cung cấp các lớp, bộ tối ưu hóa hoặc tiện ích huấn luyện ngay lập tức. Đây là một thiết kế có chủ ý. Thay vào đó, một hệ sinh thái các thư viện đã phát triển xung quanh JAX. Nổi bật nhất trong số này là Flax, một thư viện mạng nơ-ron cấp cao ban đầu được phát triển bởi nhóm Google Brain.

Flax được xây dựng xung quanh API flax.linen, cung cấp một cách có cấu trúc để định nghĩa kiến trúc mạng nơ-ron trong khi tuân thủ nghiêm ngặt triết lý hàm của JAX. Trong một mô hình Flax, các lớp không lưu trữ trọng số của riêng chúng. Thay vào đó, định nghĩa mô hình hoạt động như một bản thiết kế. Khi bạn khởi tạo một mô hình, nó trả về một từ điển lồng nhau các tham số. Trong quá trình truyền xuôi, bạn truyền rõ ràng các tham số đến phương thức apply của mô hình.

Việc quản lý trạng thái rõ ràng này ban đầu có vẻ dài dòng đối với các nhà phát triển quen thuộc với nn.Module của PyTorch. Tuy nhiên, nó tỏa sáng khi xử lý các vòng lặp tối ưu hóa phức tạp, tập hợp mô hình hoặc học siêu cấp, nơi việc có quyền truy cập rõ ràng vào toàn bộ cây tham số như một thực thể riêng biệt giúp việc thao tác trở nên cực kỳ đơn giản.

Bổ sung cho Flax là Optax, một thư viện để xử lý gradient và tối ưu hóa. Optax coi tối ưu hóa là một phép biến đổi có trạng thái được áp dụng cho các bản cập nhật. Một bộ tối ưu hóa Optax nhận các gradient và trạng thái bộ tối ưu hóa, và trả về các bản cập nhật tham số cùng với một trạng thái mới. Thiết kế này tách rời thuật toán tối ưu hóa khỏi chính mô hình.

Đi sâu: Xây dựng một Vòng lặp Huấn luyện

Để thực sự hiểu JAX, chúng ta phải xem xét cách một vòng lặp huấn luyện được cấu trúc. Bởi vì các hàm JAX phải thuần túy, chúng ta quản lý trạng thái (tham số, trạng thái bộ tối ưu hóa, khóa PRNG) từ bên ngoài.

Một bước huấn luyện điển hình trong JAX bao gồm việc gói gọn quá trình truyền xuôi, tính toán mất mát và cập nhật tối ưu hóa vào một hàm được biên dịch JIT duy nhất.

import jax
import jax.numpy as jnp
import optax
from flax.training import train_state

def create_train_state(rng, model, learning_rate):
    """Initializes the model and creates the training state."""
    params = model.init(rng, jnp.ones([1, 28, 28, 1]))['params']
    tx = optax.adam(learning_rate)
    return train_state.TrainState.create(
        apply_fn=model.apply, params=params, tx=tx)

@jax.jit
def train_step(state, batch):
    """Executes a single training step."""
    def loss_fn(params):
        logits = state.apply_fn({'params': params}, batch['image'])
        loss = optax.softmax_cross_entropy_with_integer_labels(
            logits=logits, labels=batch['label']).mean()
        return loss
    
    grad_fn = jax.value_and_grad(loss_fn)
    loss, grads = grad_fn(state.params)
    state = state.apply_gradients(grads=grads)
    return state, loss

Hãy chú ý cách hàm train_step nhận state hiện tại và trả về một state hoàn toàn mới. Bên dưới, bởi vì hàm này được biên dịch với @jax.jit, XLA tối ưu hóa các hoạt động này tại chỗ khi được thực thi trên bộ nhớ thiết bị, nghĩa là bạn có được sự thuần túy chức năng mà không phải trả giá về hiệu suất do phân bổ bộ nhớ liên tục.

Một khía cạnh quan trọng khác của JAX là bộ tạo số giả ngẫu nhiên (PRNG) rõ ràng của nó. Không giống như NumPy, có một trạng thái toàn cục ẩn cho tính ngẫu nhiên, JAX yêu cầu bạn truyền và chia khóa PRNG một cách rõ ràng bất cứ khi nào bạn cần số ngẫu nhiên. Điều này đảm bảo rằng tính ngẫu nhiên có thể tái tạo hoàn hảo và hoạt động chính xác khi các hàm được vector hóa hoặc phân tán trên nhiều thiết bị.

JAX so với các lựa chọn thay thế

Khi so sánh JAX với PyTorch hoặc TensorFlow, sự khác biệt nằm ở mức độ trừu tượng. PyTorch cung cấp một framework toàn diện, đầy đủ tính năng, trực quan cho các nhà phát triển hướng đối tượng. TensorFlow cung cấp một nền tảng đầu cuối với các công cụ triển khai mở rộng.

Ngược lại, JAX là một bộ công cụ cấp thấp hơn, thanh lịch về mặt toán học. Nó buộc nhà phát triển phải suy nghĩ theo chức năng và quản lý trạng thái một cách rõ ràng. Đối với các tác vụ học có giám sát tiêu chuẩn, PyTorch có thể nhanh hơn để thiết lập. Tuy nhiên, đối với nghiên cứu ở biên giới—chẳng hạn như kiến trúc mới lạ, mô phỏng vật lý phức tạp, học tăng cường hoặc huấn luyện phân tán quy mô lớn—các phép biến đổi có thể kết hợp của JAX mang lại sự linh hoạt và hiệu suất vượt trội. Mô hình chức năng, mặc dù có đường cong học tập, cuối cùng dẫn đến mã mạnh mẽ hơn và dễ song song hóa hơn.

Advertisement

Kết luận

Học sâu với JAX đại diện cho một sự thay đổi mạnh mẽ hướng tới sự thuần túy toán học và lập trình chức năng trong học máy. Bằng cách cung cấp các phép biến đổi có thể kết hợp như jit, grad và vmap trên các phép tính được biên dịch XLA, JAX cho phép các nhà nghiên cứu viết mã sạch sẽ có thể mở rộng dễ dàng sang phần cứng mạnh nhất thế giới. Mặc dù việc quản lý trạng thái rõ ràng của nó đòi hỏi một sự thay đổi mô hình đối với các nhà phát triển đến từ các framework hướng đối tượng truyền thống, sự rõ ràng và hiệu suất mang lại khiến JAX trở thành một công cụ không thể thiếu cho thế hệ nghiên cứu AI và điện toán hiệu suất cao tiếp theo.

Bạn cũng có thể thích

Share this article:

Stay Updated

Get the latest posts delivered straight to your inbox.

Free Developer Utilities

Free In-Browser Developer Tools

Clean AI CLI logs, build cron expressions, decode JWTs, and calculate chmod permissions offline.

Explore Tools
Advertisement
Chạy nước rút đám mây 13 ngày: Biến tín dụng GCP sắp hết hạn thành tài sản vĩnh viễn không cần bảo trì
cloud

Chạy nước rút đám mây 13 ngày: Biến tín dụng GCP sắp hết hạn thành tài sản vĩnh viễn không cần bảo trì

Hướng dẫn thực tế để tối đa hóa ROI từ các khoản tín dụng Google Cloud sắp hết hạn, giúp bạn chuyển đổi tài nguyên điện toán tạm thời thành nội dung SEO vĩnh viễn, âm thanh thần kinh và tập dữ liệu được tính toán trước với chi phí sau khi hết hạn bằng không.

Read more