از nn.Linear تا یک MLP فیوزشده در PyTorch
خلاصهٔ کاملتر
این مقاله دومین قسمت از سری «Profiling in PyTorch» از تیم Hugging Face هست. تو قسمت اول با خوندن trace پروفایلر و مفهوم سربار اجرای کرنلها آشنا شدیم؛ اینجا یه پله بالاتر میریم و به جای ضربوجمع دستی، از nn.Linear استفاده میکنیم و سهتا از این لایهها رو روی هم میچینیم تا یه بلوک MLP بسازیم. دو ایدهی پایه که مدام بهشون تکیه میشه اینه که کرنل یه برنامهست که موازی روی نخهای GPU اجرا میشه، و CPU کارش زمانبندی و لانچ این کرنلهاست؛ بیشتر سرباری که تو trace میبینی همین کار زمانبندیه.
نویسنده میگه nn.Linear در اصل همون ضرب ماتریسی بهعلاوهی bias هست. وقتی trace رو بزرگنمایی میکنی، یه op به اسم aten::t (همون transpose) قبل از aten::addmm دیده میشه. نکتهی مهم اینه که aten::t اصلاً داده رو کپی یا جابهجا نمیکنه؛ فقط متادیتای تنسور (شکل و stride) رو روی CPU بازنویسی میکنه و هیچ کرنلی روی GPU لانچ نمیشه — برای همین تو جدول، زمان CUDAش صفره.
یه چیز جالب اینه که aten::add جداگانهای برای جمع bias تو زنجیره نیست. به گفتهی نویسنده، جمع bias داخل کرنل ضرب ماتریسی فولد شده، با چیزی که بهش epilogue میگن: یه محاسبهی کوچیک که کرنل GEMM درست قبل از نوشتن نتیجه تو حافظهی HBM انجام میده. این کار باعث میشه یه بار اضافی خوندن و نوشتن تو حافظه — که گرونه — حذف بشه. در نتیجه aten::linear میبینه bias پاس داده شده و به جای دو عملیات جدا، aten::addmm رو صدا میزنه.
همینجا یه درس مهم درمیاد: واکنش رایج اینه که هر وقت مدل کند به نظر میاد سراغ torch.compile بریم. ولی برای یه GEMM تنها با bias، compile تقریباً کاری نداره چون همون کرنل cuBLAS قبلاً هم استفاده میشد؛ compile فقط چندتا op سربارِ CPU (مثل همون بازیکردن با viewها) رو حذف میکنه و GPU دقیقاً همون ریاضی رو انجام میده.
برای MLP، نویسنده یه شبکهی feed-forward با واریانت GeGLU میسازه. ساختار سادش این شکلیه:
def forward(self, x):
g = self.gate_proj(x)
u = self.up_proj(x)
h = F.gelu(g, approximate="tanh")
m = h * u
y = self.down_proj(m)
return yتو حالت eager، هر forward دقیقاً ۵ کرنل روی GPU اجرا میکنه: سه GEMM بهعلاوهی یه GeLU و یه ضرب. تنسور میانی [8192, 3072] (حدود ۵۰ مگابایت) رو کرنل GeLU تو HBM مینویسه و کرنل ضرب بلافاصله میخونتش.
نقطهی اوج درس compile همینجاست: torch.compile اون دو کرنل pointwise (GeLU و mul) بهعلاوهی یه reshape رو تو یه کرنل Triton ادغام میکنه. این کرنل g و u رو یه بار میخونه، gelu(g) * u رو حساب میکنه و نتیجه رو یه بار مینویسه؛ یه رفتوبرگشت کامل تنسور میانی از حافظهی سراسری حذف میشه. ولی سه GEMM دستنخورده میمونن و همون کرنلهای cuBLAS قبلی هستن.
در آخر نویسنده یه کرنل دستنویس و تیونشده به اسم LigerGEGLUMLP رو از کتابخونهی kernels میاره. این کرنل همون فیوژن رو بدون نیاز به کامپایلر داره و بدون Dynamo و guard و تأخیر compile کار میکنه. نکتهی صادقانهای که میگه اینه: کرنل Liger تو ۹۲.۸ میکروثانیه اجرا میشه و کرنل فیوزشدهی Inductor تو ۸۹.۴ میکروثانیه؛ پس Liger کمی کندتره. ولی Inductor برای یه شکل ثابت تخصصی شده و با تغییر batch یا seq باید دوباره trace و compile کنی، در حالی که Liger با هر شکلی بدون recompile کار میکنه. انتخاب واقعی بین «یه کرنل عمومی سریع» و «یه کرنل تخصصی برای یه شکل خاص» هست.
نکات کلیدی:
- جمع bias تو nn.Linear از قبل داخل کرنل GEMM فولد شده (epilogue)، پس یه کرنل cuBLASه نه ضرب و جمع جدا.
- aten::t و opهایی مثل reshape و view هیچ کرنلی روی GPU نمیزنن؛ فقط متادیتای تنسور رو روی CPU عوض میکنن.
- torch.compile روی یه Linear تنها چیزی برای فیوز کردن نداره؛ فقط سربار dispatch روی CPU رو کم میکنه.
- روی MLP، کرنلهای GeLU و mul تو یه کرنل Triton ادغام میشن و یه رفتوبرگشت تنسور میانی به HBM حذف میشه.
- کرنل دستنویس Liger همون فیوژن رو بدون کامپایلر میده و برخلاف compile با هر شکل ورودی بدون recompile کار میکنه.




