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

· By: SirilukP

3 Numba Tricks for Python Runtime Optimization

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 เพื่อหยุดการคอมไพล์ที่สิ้นเปลือง

Source: KDnuggets
ดูแลงานแปลและเรียบเรียงโดย SirilukP

ความคิดเห็น (0)

เข้าสู่ระบบเพื่อร่วมแสดงความเห็น

สมัครสมาชิก

มาเป็นคนแรกที่แสดงความเห็นกันเลยโบร