@@ -241,16 +241,20 @@ import torch
241241from coreai._compiler.dialects import coreai
242242from coreai_torch._utils import get_operands
243243
244+
244245@ torch.library.custom_op(" my_lib::scaled_add" , mutates_args = ())
245246def scaled_add(x: torch.Tensor, y: torch.Tensor, scale: float ) -> torch.Tensor:
246247 return x + scale * y
247248
249+
248250@ scaled_add.register_fake
249251def _(x, y, scale):
250252 return torch.empty_like(x)
251253
254+
252255converter = TorchConverter()
253256
257+
254258@ converter.register_torch_lowering(" my_lib::scaled_add.default" )
255259def lower_scaled_add(values_map, node, loc):
256260 x, y = get_operands(values_map, node, [0 , 1 ], loc)
@@ -259,6 +263,7 @@ def lower_scaled_add(values_map, node, loc):
259263 scaled_y = coreai.broadcasting_mul(y, scale_val, loc = loc)
260264 return coreai.broadcasting_add(x, scaled_y, loc = loc)
261265
266+
262267coreai_program = converter.add_exported_program(exported).to_coreai()
263268coreai_program.optimize()
264269```
@@ -272,7 +277,10 @@ from coreai_torch._utils import get_operand
272277
273278converter = TorchConverter()
274279
275- @ converter.register_torch_lowering(" aten::_adaptive_avg_pool2d.default" , allow_override = True )
280+
281+ @ converter.register_torch_lowering(
282+ " aten::_adaptive_avg_pool2d.default" , allow_override = True
283+ )
276284def lower_adaptive_avg_pool2d_static(values_map, node, loc):
277285 x = get_operand(values_map, node, 0 , loc)
278286 output_h, output_w = node.args[1 ]
@@ -290,6 +298,7 @@ def lower_adaptive_avg_pool2d_static(values_map, node, loc):
290298 coreai.cast(float (kernel_h * kernel_w), x.type.element_type),
291299 )
292300
301+
293302coreai_program = converter.add_exported_program(exported).to_coreai()
294303coreai_program.optimize()
295304```
@@ -317,7 +326,12 @@ Registers one or more `TorchMetalKernel` objects so the converter can convert th
317326
318327```python
319328import torch
320- from coreai_torch import TorchConverter, TorchMetalKernel, MetalParameter, get_decomp_table
329+ from coreai_torch import (
330+ TorchConverter,
331+ TorchMetalKernel,
332+ MetalParameter,
333+ get_decomp_table,
334+ )
321335
322336
323337def torch_add(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
@@ -383,9 +397,9 @@ coreai_program = (
383397 TorchConverter()
384398 .add_pytorch_module(
385399 model,
386- export_fn = lambda m : torch.export.export(m, args = example_inputs).run_decompositions(
387- coreai_torch.get_decomp_table()
388- ),
400+ export_fn = lambda m : torch.export.export(
401+ m, args = example_inputs
402+ ).run_decompositions(coreai_torch.get_decomp_table()) ,
389403 )
390404 .to_coreai()
391405)
@@ -457,6 +471,7 @@ class Linear(nn.Module):
457471 def forward(self , x):
458472 return self .fc(x)
459473
474+
460475ep = torch.export.export(Linear().eval(), args = (torch.randn(1 , 8 ),))
461476ep = ep.run_decompositions(get_decomp_table())
462477
@@ -473,16 +488,17 @@ TorchConverter().add_exported_program(
473488class KVCache(nn.Module):
474489 def __init__ (self ):
475490 super ().__init__ ()
476- self .register_buffer(" kv_cache" , torch.zeros(1 , 4 )) # state[0]
477- self .register_buffer(" pos_idx" , torch.zeros(1 )) # state[1]
491+ self .register_buffer(" kv_cache" , torch.zeros(1 , 4 )) # state[0]
492+ self .register_buffer(" pos_idx" , torch.zeros(1 )) # state[1]
478493
479494 def forward(self , x, y, z):
480- self .kv_cache.add_(x) # buffer mutation
481- self .pos_idx.add_(1 ) # buffer mutation
482- y.mul_(2 ) # state[2]: mutated user input
495+ self .kv_cache.add_(x) # buffer mutation
496+ self .pos_idx.add_(1 ) # buffer mutation
497+ y.mul_(2 ) # state[2]: mutated user input
483498 # non-mutated: x -> input[0], z -> input[1]
484499 return self .kv_cache + y, z * 3
485500
501+
486502ep = torch.export.export(
487503 KVCache().eval(),
488504 args = (torch.randn(1 , 4 ), torch.randn(1 , 4 ), torch.randn(1 , 4 )),
0 commit comments