forked from tinygrad/tinygrad
add var_vals to kopt with symbolic (#2008)
* add var_vals to kopt with symbolic again * no copies
This commit is contained in:
@@ -4,8 +4,7 @@ from tinygrad.tensor import Tensor
|
||||
from tinygrad.nn import optim
|
||||
from tinygrad.nn.state import get_parameters
|
||||
from tinygrad.jit import TinyJit, JIT_SUPPORTED_DEVICE
|
||||
from tinygrad.ops import GlobalCounters, LazyOp, LoadOps
|
||||
from tinygrad.ops import Device
|
||||
from tinygrad.ops import Device, GlobalCounters, LazyOp, LoadOps
|
||||
from tinygrad.helpers import CI, dtypes, getenv, prod
|
||||
from tinygrad.codegen.search import kernel_optimize_opts
|
||||
|
||||
@@ -14,7 +13,7 @@ from examples.hlb_cifar10 import SpeedyResNet
|
||||
from examples.llama import Transformer as LLaMaTransformer, MODEL_PARAMS as LLAMA_MODEL_PARAMS
|
||||
from examples.stable_diffusion import UNetModel
|
||||
|
||||
def kopt_search_hook(k, create_k, to_prg, baseline, bufs):
|
||||
def kopt_search_hook(k, create_k, to_prg, baseline, bufs, var_vals):
|
||||
import nevergrad as ng
|
||||
wanna_output = bufs[0].toCPU().copy()
|
||||
def check_opt(x):
|
||||
@@ -22,7 +21,7 @@ def kopt_search_hook(k, create_k, to_prg, baseline, bufs):
|
||||
k = create_k()
|
||||
k.apply_auto_opt(x)
|
||||
prg = to_prg(k)
|
||||
first_tm = prg.exec(bufs, force_wait=True, optimizing=True)
|
||||
first_tm = prg.exec(bufs, var_vals, force_wait=True, optimizing=True)
|
||||
np.testing.assert_allclose(wanna_output, bufs[0].toCPU(), atol=1e-4, rtol=1e-4)
|
||||
return first_tm
|
||||
except Exception:
|
||||
|
||||
@@ -2,6 +2,7 @@ from typing import Callable
|
||||
import time
|
||||
from tinygrad.codegen.linearizer import Linearizer
|
||||
from tinygrad.helpers import DEBUG, prod, getenv
|
||||
from tinygrad.lazy import var_vals_from_ast
|
||||
|
||||
def get_divisors(n, min_div = 1, max_div = 512):
|
||||
if min_div > 1: yield 1
|
||||
@@ -21,16 +22,16 @@ def kernel_optimize_opts(k:Linearizer):
|
||||
opts.append(ng.p.TransitionChoice([(i,s,"G") for s in get_divisors(k.full_shape[k.first_reduce+i], min_div=4) if all(st.shape[k.first_reduce+i] % s == 0 or st.shape[k.first_reduce+i] == 1 for st in k.sts)]))
|
||||
return opts
|
||||
|
||||
def kernel_optimize_search(k:Linearizer, create_k:Callable[[], Linearizer], to_prg, baseline, bufs):
|
||||
def kernel_optimize_search(k:Linearizer, create_k:Callable[[], Linearizer], to_prg, baseline, bufs, var_vals):
|
||||
import nevergrad as ng
|
||||
def opt(x):
|
||||
try:
|
||||
k = create_k()
|
||||
k.apply_auto_opt(x)
|
||||
prg = to_prg(k)
|
||||
first_tm = prg.exec(bufs, force_wait=True, optimizing=True)
|
||||
first_tm = prg.exec(bufs, var_vals, force_wait=True, optimizing=True)
|
||||
if baseline*5 < first_tm*1000: return first_tm*1000 # very slow
|
||||
tm = min([first_tm]+[prg.exec(bufs, force_wait=True, optimizing=True) for _ in range(2)])*1000
|
||||
tm = min([first_tm]+[prg.exec(bufs, var_vals, force_wait=True, optimizing=True) for _ in range(2)])*1000
|
||||
return tm
|
||||
except Exception:
|
||||
if DEBUG >= 3:
|
||||
@@ -65,13 +66,14 @@ def kernel_optimize(k:Linearizer, create_k:Callable[[], Linearizer], to_prg, buf
|
||||
# don't optimize variable shapes
|
||||
choice = "BASELINE"
|
||||
else:
|
||||
var_vals = {k:k.min for k in var_vals_from_ast(k.ast)}
|
||||
# get baseline
|
||||
def get_baseline():
|
||||
k = create_k()
|
||||
k.hand_coded_optimizations()
|
||||
prg = to_prg(k)
|
||||
return min([prg.exec(bufs, force_wait=True, optimizing=True) for _ in range(5)])*1000
|
||||
choice = kernel_optimize_search(k, create_k, to_prg, get_baseline(), bufs)
|
||||
return min([prg.exec(bufs, var_vals, force_wait=True, optimizing=True) for _ in range(5)])*1000
|
||||
choice = kernel_optimize_search(k, create_k, to_prg, get_baseline(), bufs, var_vals)
|
||||
if global_db is not None:
|
||||
global_db[skey] = choice
|
||||
global_db.sync()
|
||||
|
||||
Reference in New Issue
Block a user