Demo: nhìn AI học
Thuộc loạt AI cho kỹ sư nhúng. Số thật, sinh bằng
src/trace_learn.py.
Mạng nơ-ron thật có hàng triệu tham số, không nhìn được gì. Mạng dưới đây có đúng 9 tham số, đủ nhỏ để in hết lên màn hình và xem từng con số thay đổi:
x[2] --> Linear(2->2) --> ReLU --> Linear(2->1) --> y
W1(4) b1(2) W2(2) b2(1) = 9 tham số
Nhiệm vụ: cho x = [0.8, -0.5], đẩy đầu ra y về 1.0. Kéo thanh trượt để đi qua từng bước huấn luyện.
Điều đáng nhìn nhất
Kéo thanh trượt từ bước 0 tới bước 39 và để ý cột autograd so với tự tính tay ở khối 3. Hai cột đó là hai phép tính hoàn toàn độc lập: một bên là loss.backward() của PyTorch, một bên là công thức đạo hàm viết tay trong trace_learn.py. Chúng khớp nhau tới chữ số cuối, sai lệch lớn nhất qua 40 bước là 0.0.
Đó là toàn bộ nội dung của backpropagation. Nó không phải một thuật toán bí ẩn, nó là chain rule áp dụng ngược đồ thị tính toán, và bạn kiểm lại bằng tay được.
Thử ReLU chết
Đổi seed là gặp ngay một hiện tượng đáng nhớ:
cd src && uv run python trace_learn.py --seed 7 # nơ-ron 0 chết
cd src && uv run python trace_learn.py --seed 5 # CẢ HAI nơ-ron chết
Với --seed 7, z[0] = -0.4084 ở bước 0 và không nhúc nhích suốt 40 bước. Lý do nằm ở đúng một dòng trong chain rule:
dL/dz = dL/dh * [z > 0]
z[0] âm nên [z > 0] bằng 0, nên gradient chảy về nhánh đó bằng 0, nên trọng số không được cập nhật, nên z[0] vẫn âm ở bước sau. Nơ-ron đó chết vĩnh viễn ngay từ lúc khởi tạo. Với --seed 5 thì cả hai chết, và mạng chỉ còn học được qua b2.
Đây là lý do khởi tạo weight quan trọng, và là chỗ nối thẳng sang bài 2, mục 2.4.
Bài viết thuộc loạt AI cho kỹ sư nhúng. Góp ý: mở issue tại github.com/ninhnn2/machineai.