3 เทคนิคใช้ Numba เพิ่มความเร็ว Python Runtime ให้แรงกว่าเดิม

Numba คือตัวช่วยคอมไพล์ loop ของ Python เชิงตัวเลขให้กลายเป็น machine code โดยที่คุณไม่ต้องออกจากสภาพแวดล้อม Python หรือเขียนโค้ดใหม่ด้วยภาษา C แม้แต่การทำ vectorize ในส่วนที่ยุ่งยากก็ไม่ใช่ปัญหา
เมื่อผลลัพธ์ของ Numba ไม่เป็นไปตามคาด ปัญหามักไม่ได้อยู่ที่ตัวคอมไพเลอร์ แต่เป็นเรื่องของขอบเขตรอบโค้ดที่ถูกคอมไพล์ ไม่ว่าจะเป็นการไม่ข้ามขอบเขตนั้น การขยายขอบเขตไม่กว้างพอ หรือการข้ามขอบเขตเดิมซ้ำๆ ทุกครั้งที่รัน ต่อไปนี้คือ 3 เทคนิคที่จะช่วยแก้ปัญหาดังกล่าว ซึ่งผ่านการตรวจสอบกับ Numba 0.67.0 เรียบร้อยแล้ว
pip install numbaเทคนิคที่ 1: การคอมไพล์ Loop แทนที่จะเป็นการ Interpret
ตามปกติแล้วการทำ reduction บน NumPy array จะช้าเพราะ Python interpreter ต้องส่งต่อข้อมูล (dispatch) ตามชนิดข้อมูลทีละองค์ประกอบหลายล้านครั้ง แต่ด้วยการเพิ่ม decorator เพียงบรรทัดเดียว Numba จะอ่านชนิดข้อมูลในการเรียกใช้ total_jit() ครั้งแรกเพื่อคอมไพล์ให้เป็นแบบเฉพาะทาง หลังจากนั้นการเรียกใช้ทั้งหมดจะเป็นการรัน native code โดยตรง
import time
import numpy as np
from numba import njit
def total_plain(x):
total = 0.0
for i in range(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
@njit
def total_jit(x):
total = 0.0
for i in range(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return total
def total_numpy(x):
return np.sum(np.sqrt(x) * np.sin(x))
def benchmark(func, x, repeats=3):
"""Run func repeats times and return (best_time, mean_time, result)."""
times = []
result = None
for _ in range(repeats):
start = time.perf_counter()
result = func(x)
times.append(time.perf_counter() - start)
return min(times), sum(times) / len(times), result
x = np.random.default_rng(0).random(10_000_000)
# Measure JIT compilation separately (first call compiles)
start = time.perf_counter()
total_jit(x)
compile_time = time.perf_counter() - start
print(f"Numba first call (includes compilation): {compile_time:.4f} s\n")
# Plain Python loop is slow on 10M elements, so run it only once
results = {
"Plain Python loop": benchmark(total_plain, x, repeats=1),
"Numba @njit": benchmark(total_jit, x, repeats=5),
"NumPy vectorized": benchmark(total_numpy, x, repeats=5),
}
baseline = results["Plain Python loop"][0]
print(f"{'Method':<20} {'Best (s)':>10} {'Mean (s)':>10} {'Speedup':>10} Result")
print("-" * 72)
for name, (best, mean, value) in results.items():
print(f"{name:<20} {best:>10.4f} {mean:>10.4f} {baseline / best:>9.1f}x {value:.6f}")
# Sanity check that all methods agree
values = [r[2] for r in results.values()]
print("\nResults match:", np.allclose(values, values[0]))ผลลัพธ์ที่ได้แสดงให้เห็นว่าความเร็วเพิ่มขึ้นอย่างมาก:
Numba first call (includes compilation): 0.2799 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0789 3.0789 1.0x 3641603.675817
Numba @njit 0.0384 0.0385 80.2x 3641603.675817
NumPy vectorized 0.0552 0.0606 55.7x 3641603.675816
Results match: Trueอย่างไรก็ตาม การใช้ Nopython mode (@njit) มีข้อจำกัดที่สำคัญคือคุณต้องรักษาฟังก์ชันที่ทำงานหนัก (hot function) ให้อยู่ในขอบเขตของ Python และ NumPy ที่ Numba สามารถระบุชนิดข้อมูลได้ชัดเจนเท่านั้น
เทคนิคที่ 2: การกระจาย Loop ไปยังทุก Core
แม้ฟังก์ชันจะถูกคอมไพล์แล้วแต่ก็ยังรันบนคอร์เดียว การเพิ่ม parallel=True จะทำให้ Numba พยายามประมวลผลแบบขนาน และการเปลี่ยนมาใช้ prange() จะเป็นการระบุว่าต้องการรัน loop ไหนแบบขนาน โดยแทบไม่ต้องเปลี่ยนโค้ดส่วนอื่นเลย
from numba import njit, prange
@njit(parallel=True)
def total_parallel(x):
total = 0.0
for i in prange(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return totalเมื่อรันการทดสอบเปรียบเทียบใหม่ ผลลัพธ์ที่ได้คือความเร็วที่เพิ่มขึ้นอย่างก้าวกระโดด:
Numba first call (includes compilation): 0.2431 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0229 3.0229 1.0x 3641603.675817
Numba @njit 0.0374 0.0379 80.7x 3641603.675817
NumPy vectorized 0.0547 0.0608 55.3x 3641603.675816
Parallel @njit 0.0087 0.1222 347.3x 3641603.675816
Results match: Trueสาเหตุที่ทำเช่นนี้ได้ปลอดภัยเพราะ Numba จดจำรูปแบบ total += ... เป็นการทำ reduction โดยจะแบ่งช่วงการทำงานกระจายไปยัง thread ต่างๆ และนำผลลัพธ์มารวมกันในตอนท้าย ซึ่งใช้ได้กับเครื่องหมายทางคณิตศาสตร์อื่นๆ เช่น -=, *=, /= รวมถึงฟังก์ชัน max และ min ด้วย
เทคนิคที่ 3: จ่ายค่าคอมไพล์เพียงครั้งเดียว เนื่องจากการคอมไพล์จะเกิดขึ้นในการเรียกใช้ครั้งแรก หากคุณต้องรันสคริปต์บ่อยๆ runtime ส่วนใหญ่จะถูกใช้ไปกับการคอมไพล์ซ้ำ การใช้ cache=True จะช่วยบันทึกผลการคอมไพล์ลงในดิสก์ ทำให้การรันครั้งต่อไปสามารถโหลดมาใช้งานได้ทันทีโดยไม่ต้องคอมไพล์ใหม่
@njit(parallel=True, cache=True)
def total_cached(x):
total = 0.0
for i in prange(x.shape[0]):
total += np.sqrt(x[i]) * np.sin(x[i])
return totalผลลัพธ์จากการใช้ Cache:
Numba first call (includes compilation): 0.1248 s
Method Best (s) Mean (s) Speedup Result
------------------------------------------------------------------------
Plain Python loop 3.0025 3.0025 1.0x 3641603.675817
Numba @njit 0.0384 0.0384 78.2x 3641603.675817
NumPy vectorized 0.0550 0.0563 54.5x 3641603.675816
Parallel @njit 0.0086 0.0383 349.1x 3641603.675816
Cached @njit 0.0086 0.0094 349.9x 3641603.675816
Results match: Trueข้อควรระวังคือตัวแปร Global จะถูกสตัฟฟ์ (frozen) ไว้ตามค่าเดิมตอนคอมไพล์ และระบบ cache อาจไม่รับรู้ถึงการเปลี่ยนแปลงในไฟล์อื่นที่นำมาใช้เป็นฟังก์ชันตัวช่วย นอกจากนี้การทำ caching ร่วมกับ parallel=True ยังมีความซับซ้อน จึงควรตรวจสอบให้แน่ใจว่าระบบดึงข้อมูลจาก cache จริงก่อนใช้งาน
สรุป
การเพิ่มประสิทธิภาพด้วย Numba สรุปได้สั้นๆ คือการตั้งคำถามว่า งานอยู่ในขอบเขตคอมไพล์หรือไม่, ใช้ทรัพยากรเครื่องได้เต็มที่หรือไม่ และคุณจ่ายต้นทุนการคอมไพล์ซ้ำซ้อนเกินไปหรือไม่ เพียงเริ่มจากคอมไพล์ loop ขยายการทำงานให้ขนานกัน และปิดท้ายด้วยการทำ cache เพื่อหยุดการคอมไพล์ที่สิ้นเปลือง
ความคิดเห็น (0)
เข้าสู่ระบบเพื่อร่วมแสดงความเห็น
สมัครสมาชิกมาเป็นคนแรกที่แสดงความเห็นกันเลยโบร
