add var_vals to kopt with symbolic (#2008)

* add var_vals to kopt with symbolic again

* no copies
This commit is contained in:
nimlgen
2023-10-07 09:34:21 -07:00
committed by GitHub
parent 121f7aa8c5
commit d07ac379f9
2 changed files with 10 additions and 9 deletions
+3 -4
View File
@@ -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:
+7 -5
View File
@@ -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()