forked from tinygrad/tinygrad
use the device abstraction in handcode_resnet50_opt
This commit is contained in:
@@ -1,10 +1,9 @@
|
||||
from typing import List
|
||||
from models.resnet import ResNet50
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.ops import LoadOps
|
||||
from tinygrad.ops import LoadOps, Device, Compiled
|
||||
from tinygrad.codegen.kernel import LinearizerOptions
|
||||
from tinygrad.codegen.linearizer import Linearizer
|
||||
from tinygrad.runtime.ops_metal import renderer, MetalProgram, RawMetalBuffer
|
||||
from tinygrad.helpers import ansilen, DEBUG
|
||||
from tinygrad.graph import print_tree
|
||||
|
||||
@@ -12,6 +11,10 @@ if __name__ == "__main__":
|
||||
mdl = ResNet50()
|
||||
seen = set()
|
||||
|
||||
# the device we are optimizing for
|
||||
device: Compiled = Device[Device.DEFAULT]
|
||||
print(f"optimizing for {Device.DEFAULT}")
|
||||
|
||||
# first model run to init the weights, they are saved in seen
|
||||
mdl(Tensor.empty(64, 3, 224, 224)).lazydata.schedule(seen)
|
||||
|
||||
@@ -32,12 +35,12 @@ if __name__ == "__main__":
|
||||
# "linearize" the op into uops in different ways
|
||||
lins:List[Linearizer] = []
|
||||
|
||||
if i == 1:
|
||||
if Device.DEFAULT == "METAL" and i == 1:
|
||||
# through careful work, we discovered 1,8,0
|
||||
for big_chomp in [1,2]: #[1,2,4,8,16]:
|
||||
for lil_chomp in [2,4,7,8,14]:
|
||||
for upcasted in [0,1,2]:
|
||||
lin = Linearizer(si.ast, LinearizerOptions(device="METAL"))
|
||||
lin = Linearizer(si.ast, device.linearizer_opts)
|
||||
lin.reshape_and_permute(lambda x: (4096//big_chomp,big_chomp,56//lil_chomp,lil_chomp,56//lil_chomp,lil_chomp)+x[-2:], [0,2,4,1,3,5,6,7])
|
||||
lin.upcasted += upcasted
|
||||
lin.local_dims += 3
|
||||
@@ -45,28 +48,21 @@ if __name__ == "__main__":
|
||||
else:
|
||||
# try with and without tensor cores
|
||||
for tc in [0,1]:
|
||||
lin = Linearizer(si.ast, LinearizerOptions(device="METAL"))
|
||||
lin = Linearizer(si.ast, device.linearizer_opts)
|
||||
lin.hand_coded_optimizations(use_tensor_cores=tc)
|
||||
lins.append(lin)
|
||||
|
||||
# create output/input buffers
|
||||
rout = RawMetalBuffer(si.out.st.size(), si.out.dtype)
|
||||
rin = [RawMetalBuffer(x.st.size(), x.dtype) for x in si.inputs]
|
||||
rawbufs = [device.buffer(si.out.st.size(), si.out.dtype)] + [device.buffer(x.st.size(), x.dtype) for x in si.inputs]
|
||||
|
||||
# benchmark the programs
|
||||
choices = []
|
||||
for lin in lins:
|
||||
# render the code and create the program
|
||||
lin.linearize()
|
||||
code = renderer(lin.function_name, lin.uops)
|
||||
prg = MetalProgram(lin.function_name, code)
|
||||
|
||||
# print the kernel code if you want
|
||||
#print(code)
|
||||
prg = device.to_program(lin)
|
||||
|
||||
# benchmark it by running 10 times
|
||||
try:
|
||||
tm = min([prg(lin.global_size, lin.local_size, rout, *rin, wait=True) for _ in range(10)])
|
||||
tm = min([prg(rawbufs, force_wait=True) for _ in range(10)])
|
||||
choices.append((tm, lin))
|
||||
except AssertionError:
|
||||
tm = float('inf')
|
||||
|
||||
Reference in New Issue
Block a user