From d07ac379f95e04c56e9fd99ece6c8545a4fd8a85 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Sat, 7 Oct 2023 19:34:21 +0300 Subject: [PATCH] add var_vals to kopt with symbolic (#2008) * add var_vals to kopt with symbolic again * no copies --- test/models/test_real_world.py | 7 +++---- tinygrad/codegen/search.py | 12 +++++++----- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/test/models/test_real_world.py b/test/models/test_real_world.py index 4e2f8f0466..1e28fb82ce 100644 --- a/test/models/test_real_world.py +++ b/test/models/test_real_world.py @@ -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: diff --git a/tinygrad/codegen/search.py b/tinygrad/codegen/search.py index c727275e02..33786b2ae1 100644 --- a/tinygrad/codegen/search.py +++ b/tinygrad/codegen/search.py @@ -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()