-
Notifications
You must be signed in to change notification settings - Fork 48
add support for Operators for Generic target needed in MAGIA (again) #195
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: devel
Are you sure you want to change the base?
Changes from all commits
025e893
142b77b
bbb6fa6
7fa109b
bb9a623
67ef646
d02e0f4
fe21445
41fc725
e20401f
35f3901
9e495bf
3514c22
7ee2e74
f646dfd
36e3485
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -1311,6 +1311,10 @@ def typeCheckNodeInputs(self, ctxt: NetworkContext, node: gs.Node) -> bool: | |
| if not isinstance(reference, VariableBuffer): | ||
| return False | ||
|
|
||
| # Absent optional input: zero-sized placeholder, nothing to type check | ||
| if hasattr(reference, "values") and reference.values.size == 0: | ||
| continue | ||
|
|
||
| if hasattr(reference, "values"): | ||
| retCheck &= _type.referencedType.checkPromotion(reference.values) | ||
| else: | ||
|
|
@@ -1328,6 +1332,11 @@ def typeInferGlobalCtxt(self, ctxt: NetworkContext, node: gs.Node) -> NetworkCon | |
| for inputNode, _type in zip(node.inputs, self.input_types): | ||
| if isinstance(ctxt.lookup(inputNode.name), ConstantBuffer): | ||
| reference = ctxt.lookup(inputNode.name) | ||
|
|
||
| # Absent optional input: zero-sized placeholder, nothing to infer | ||
| if reference.values.size == 0: | ||
| continue | ||
|
|
||
| if not _type.referencedType.checkPromotion(reference.values): | ||
| raise Exception(f"Can't cast {reference} to {_type}!") | ||
|
|
||
|
|
@@ -1914,7 +1923,8 @@ def broadcast(self, ctxt: NetworkContext, default_channels_first: bool = True) - | |
| newInputShapes, newOutputShapes = self.computeShapes(inputShapes, outputShapes, | ||
| self.mapper.parser.operatorRepresentation, channels_first) | ||
|
|
||
| for node, newShape in zip(self.node.inputs + self.node.outputs, newInputShapes + newOutputShapes): | ||
| for node, newShape in zip(self.node.inputs + self.node.outputs, newInputShapes + newOutputShapes, | ||
| strict = True): | ||
| if ctxt.is_local(node.name): | ||
| ctxt.localObjects[node.name].shape = newShape | ||
| # Update shape of tensors in onnx graph | ||
|
|
@@ -2103,7 +2113,7 @@ def bind(self, ctxt: NetworkContext) -> Tuple[NetworkContext, bool]: | |
| npType = self._broadcastToNpType(ctxt.localObjects[node.name]._type) | ||
| if npType is not None: | ||
| node.dtype = npType | ||
| elif ctxt.is_global(node.name): | ||
| elif ctxt.is_global(node.name) and hasattr(ctxt.globalObjects[node.name], '_type'): | ||
| npType = self._broadcastToNpType(ctxt.globalObjects[node.name]._type) | ||
| if isinstance(ctxt.globalObjects[node.name], ConstantBuffer): | ||
| if isinstance(node, gs.Constant): | ||
|
|
@@ -2961,6 +2971,8 @@ def generateBufferInitializationCode(self) -> str: | |
| callStack = '' | ||
| for node in ctxt.globalObjects.values(): | ||
| if isinstance(node, VariableBuffer) and not isinstance(node, StructBuffer): | ||
| if not hasattr(node, '_type'): | ||
| continue | ||
|
Comment on lines
+2974
to
+2975
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🩺 Stability & Availability | 🟠 Major | ⚡ Quick win Apply the untyped-buffer guard to deallocation. These branches skip initialization and allocation for untyped global buffers. Add the same guard to deallocation, or mark optional-input placeholders as non-deployable. Also applies to: 3021-3022 🤖 Prompt for AI Agents |
||
| assert issubclass(node._type, Pointer), f"Global VariableBuffer {node.name} is not a Pointer!" | ||
| if node._deploy: | ||
| name = node.name | ||
|
|
@@ -3006,6 +3018,8 @@ def generateBufferAllocationCode(self) -> str: | |
|
|
||
| for node in ctxt.globalObjects.values(): | ||
| if isinstance(node, VariableBuffer) and not isinstance(node, StructBuffer): | ||
| if not hasattr(node, '_type'): | ||
| continue | ||
| assert issubclass(node._type, Pointer), f"Global VariableBuffer {node.name} is not a Pointer!" | ||
| if node._deploy: | ||
| name = node.name | ||
|
|
@@ -3395,6 +3409,29 @@ def _mangleNodeNames(self): | |
| seen[orig] = idx + 1 | ||
| # else: unique name, leave it unchanged | ||
|
|
||
| # Don't override this | ||
| def _nameEmptyTensors(self): | ||
| """Assign a name to every unnamed tensor in the graph | ||
|
|
||
| Deeploy keys every tensor by its name, so this pass replaces each | ||
| unnamed input with a uniquely named, zero-sized Constant. | ||
| """ | ||
| takenNames = set(self.graph.tensors().keys()) | ||
|
|
||
| for node in self.graph.nodes: | ||
| for idx, tensor in enumerate(node.inputs): | ||
| if not tensor.is_empty(): | ||
| continue | ||
|
|
||
| baseName = f"{node.name or node.op}_empty_input_{idx}" | ||
| name, counter = baseName, 0 | ||
| while name in takenNames: | ||
| counter += 1 | ||
| name = f"{baseName}_{counter}" | ||
| takenNames.add(name) | ||
|
|
||
| node.inputs[idx] = gs.Constant(name, np.zeros(0, dtype = np.float32)) | ||
|
coderabbitai[bot] marked this conversation as resolved.
|
||
|
|
||
| # Don't override this | ||
| def _removeIdentityNodes(self): | ||
| for node in filter(lambda x: x.op == "Identity", self.graph.nodes): | ||
|
|
@@ -3419,6 +3456,9 @@ def frontEnd(self): | |
| log.debug(" - Remove Identity Nodes") | ||
| self._removeIdentityNodes() | ||
|
|
||
| log.debug(" - Name Empty Tensors") | ||
| self._nameEmptyTensors() | ||
|
|
||
| log.debug(" - Mangle Tensor Names") | ||
| self._mangleTensorNames() | ||
|
|
||
|
|
@@ -3542,6 +3582,8 @@ def _printMemorySummary(self): | |
| # We do not count structs for now, since they are not properly modeled | ||
| if isinstance(_buffer, ConstantBuffer) or (isinstance(_buffer, VariableBuffer) and _buffer._deploy): | ||
| # SCHEREMO: We only | ||
| if not hasattr(_buffer, '_type'): | ||
| continue | ||
| if (hasattr(_buffer, "_memoryLevel") and _buffer._memoryLevel == level) or level == "None": | ||
| staticSize += int((np.prod(_buffer.shape) * _buffer._type.referencedType.typeWidth // 8)) | ||
| else: | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -11,23 +11,25 @@ | |
| int8_t, int32_t, uint8_t | ||
| from Deeploy.DeeployTypes import CodeTransformation, NodeBinding | ||
| from Deeploy.FutureExtension.CodeTransformationPasses.FutureCodeTransformation import FutureGeneration | ||
| from Deeploy.Targets.Generic.Templates import AddTemplate, BatchNormalizationTemplate, ConcatTemplate, ConvTemplate, \ | ||
| ConvTransposeTemplate, DebugPrintTemplate, DequantTemplate, DummyTemplate, DWConvTemplate, FloatAddTemplate, \ | ||
| FloatAveragePoolTemplate, FloatCeilTemplate, FloatClipTemplate, FloatConvTemplate, FloatDivTemplate, \ | ||
| FloatDWConvTemplate, FloatExpTemplate, FloatFloorTemplate, FloatGELUTemplate, FloatGemmTemplate, \ | ||
| FloatGlobalAveragePoolTemplate, FloatGlobalMaxPoolTemplate, FloatGroupNormTemplate, FloatHardSigmoidTemplate, \ | ||
| FloatHardSwishTemplate, FloatInstanceNormTemplate, FloatLayernormTemplate, FloatMatMulTemplate, \ | ||
| FloatMaxPoolTemplate, FloatMulTemplate, FloatPadTemplate, FloatPowTemplate, FloatReduceMeanTemplate, \ | ||
| FloatReluTemplate, FloatSigmoidTemplate, FloatSoftmaxTemplate, FloatSqrtTemplate, FloatSubTemplate, \ | ||
| FloatSwishTemplate, GatherTemplate, GemmTemplate, IntegerDivTemplate, ITAMaxTemplate, ITAPartialMaxTemplate, \ | ||
| MatMulTemplate, MaxPoolTemplate, MulTemplate, PadTemplate, QuantTemplate, ReduceMeanTemplate, ReduceSumTemplate, \ | ||
| RequantShiftTemplate, ReshapeTemplate, RQIntegerDivTemplate, RQSiGELUTemplate, SliceTemplate, SubTemplate, \ | ||
| TransposeTemplate, iGELUTemplate, iLayernormTemplate, iRMSNormTemplate, iSoftmaxTemplate | ||
| from Deeploy.Targets.Generic.Templates import AddTemplate, BatchNormalizationTemplate, Col2ImTemplate, ConcatTemplate, \ | ||
| ConvTemplate, ConvTransposeTemplate, DebugPrintTemplate, DequantTemplate, DummyTemplate, DWConvTemplate, \ | ||
| FloatAddTemplate, FloatAveragePoolTemplate, FloatCeilTemplate, FloatClipTemplate, FloatConvTemplate, \ | ||
| FloatDivTemplate, FloatDWConvTemplate, FloatEluTemplate, FloatExpTemplate, FloatFloorTemplate, FloatGELUTemplate, \ | ||
| FloatGemmTemplate, FloatGlobalAveragePoolTemplate, FloatGlobalMaxPoolTemplate, FloatGroupNormTemplate, \ | ||
| FloatHardSigmoidTemplate, FloatHardSwishTemplate, FloatInstanceNormTemplate, FloatLayernormTemplate, \ | ||
| FloatLeakyReluTemplate, FloatMatMulTemplate, FloatMaxPoolTemplate, FloatMulTemplate, FloatPadTemplate, \ | ||
| FloatPowTemplate, FloatReduceMeanTemplate, FloatReluTemplate, FloatSeluTemplate, FloatSigmoidTemplate, \ | ||
| FloatSoftmaxTemplate, FloatSqrtTemplate, FloatSubTemplate, FloatSwishTemplate, FloatTanhTemplate, GatherTemplate, \ | ||
| GemmTemplate, IntegerDivTemplate, ITAMaxTemplate, ITAPartialMaxTemplate, MatMulTemplate, MaxPoolTemplate, \ | ||
| MulTemplate, PadTemplate, QuantTemplate, ReduceMeanTemplate, ReduceSumTemplate, RequantShiftTemplate, \ | ||
| ReshapeTemplate, ResizeTemplate, RQIntegerDivTemplate, RQSiGELUTemplate, ScatterTemplate, SliceTemplate, \ | ||
| SplitTemplate, SubTemplate, TransposeTemplate, iGELUTemplate, iLayernormTemplate, iRMSNormTemplate, \ | ||
| iSoftmaxTemplate | ||
| from Deeploy.Targets.Generic.TypeCheckers import AddChecker, BatchNormChecker, ConcatChecker, ConvChecker, \ | ||
| DebugPrintChecker, DequantChecker, DivChecker, DummyChecker, GatherChecker, GELUChecker, GEMMChecker, \ | ||
| LayerNormChecker, MatMulChecker, MaxPoolChecker, MulChecker, PadChecker, QuantChecker, ReduceMeanChecker, \ | ||
| ReduceSumChecker, ReluChecker, RequantShiftChecker, ReshapeChecker, RQIntegerDivChecker, SliceChecker, \ | ||
| SoftmaxChecker, TransposeChecker | ||
| LayerNormChecker, MatMulChecker, MaxPoolChecker, MulChecker, PadChecker, PassThroughTypeChecker, QuantChecker, \ | ||
| ReduceMeanChecker, ReduceSumChecker, ReluChecker, RequantShiftChecker, ReshapeChecker, RQIntegerDivChecker, \ | ||
| SigmoidChecker, SliceChecker, SoftmaxChecker, SplitChecker, TransposeChecker | ||
|
|
||
| BasicTransformer = CodeTransformation([ArgumentStructGeneration(), MemoryManagementGeneration(), FutureGeneration()]) | ||
|
|
||
|
|
@@ -305,6 +307,11 @@ | |
| ConcatTemplate.referenceTemplate, BasicTransformer) | ||
| ] | ||
|
|
||
| BasicSplitBindings = [ | ||
| NodeBinding(SplitChecker([PointerClass(type), PointerClass(int32_t)], [PointerClass(type)]), | ||
| SplitTemplate.referenceTemplate, BasicTransformer) for type in IntegerDataTypes + FloatDataTypes | ||
| ] | ||
|
|
||
| BasicQuantBindings = [ | ||
| NodeBinding(QuantChecker([PointerClass(float32_t)], [PointerClass(int8_t)]), QuantTemplate.referenceTemplate, | ||
| BasicTransformer), | ||
|
|
@@ -329,19 +336,35 @@ | |
| for type in FloatDataTypes | ||
| ] | ||
|
|
||
| BasicConvTransposeBindings = [ | ||
| BasicConvTranspose1DBindings = [ | ||
| NodeBinding( | ||
| ConvChecker( | ||
| [PointerClass(dtype), PointerClass(dtype), PointerClass(dtype)], # input, weight, bias | ||
| [PointerClass(dtype)]), | ||
| ConvTransposeTemplate.referenceTemplate1D, | ||
| BasicTransformer) for dtype in FloatDataTypes | ||
| ] + [ | ||
| NodeBinding( | ||
| ConvChecker( | ||
| [PointerClass(dtype), PointerClass(dtype)], # input, weight | ||
| [PointerClass(dtype)]), | ||
| ConvTransposeTemplate.referenceTemplate1D, | ||
| BasicTransformer) for dtype in FloatDataTypes | ||
| ] | ||
|
|
||
| BasicConvTranspose2DBindings = [ | ||
| NodeBinding( | ||
| ConvChecker( | ||
| [PointerClass(type), PointerClass(type), PointerClass(type)], # input, weight, bias | ||
| [PointerClass(type)]), | ||
| ConvTransposeTemplate.referenceTemplate, | ||
| ConvTransposeTemplate.referenceTemplate2D, | ||
| BasicTransformer) for type in FloatDataTypes | ||
| ] + [ | ||
| NodeBinding( | ||
| ConvChecker( | ||
| [PointerClass(type), PointerClass(type)], # input, weight | ||
| [PointerClass(type)]), | ||
| ConvTransposeTemplate.referenceTemplate, | ||
| ConvTransposeTemplate.referenceTemplate2D, | ||
| BasicTransformer) for type in FloatDataTypes | ||
|
Comment on lines
+339
to
368
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | 🏗️ Heavy lift Do not bind unsupported ConvTranspose attributes.
Propagate these attributes through the templates and kernel interfaces. Otherwise, reject every unsupported non-default value in 🧰 Tools🪛 Ruff (0.16.1)[error] 355-355: Variable (A001) [error] 362-362: Variable (A001) 🤖 Prompt for AI Agents |
||
| ] | ||
|
|
||
|
|
@@ -368,8 +391,13 @@ | |
| BasicTransformer), | ||
| ] | ||
|
|
||
| BasicTanhBindings = [ | ||
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), FloatTanhTemplate.referenceTemplate, | ||
| BasicTransformer), | ||
| ] | ||
|
|
||
| BasicSigmoidBindings = [ | ||
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), | ||
| NodeBinding(SigmoidChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), | ||
| FloatSigmoidTemplate.referenceTemplate, BasicTransformer), | ||
| ] | ||
|
|
||
|
|
@@ -388,6 +416,21 @@ | |
| FloatHardSwishTemplate.referenceTemplate, BasicTransformer), | ||
| ] | ||
|
|
||
| BasicEluBindings = [ | ||
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), FloatEluTemplate.referenceTemplate, | ||
| BasicTransformer), | ||
| ] | ||
|
|
||
| BasicSeluBindings = [ | ||
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), FloatSeluTemplate.referenceTemplate, | ||
| BasicTransformer), | ||
| ] | ||
|
|
||
| BasicLeakyReluBindings = [ | ||
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), | ||
| FloatLeakyReluTemplate.referenceTemplate, BasicTransformer), | ||
| ] | ||
|
|
||
| BasicInstanceNormBindings = [ | ||
| NodeBinding( | ||
| DummyChecker( | ||
|
|
@@ -423,3 +466,22 @@ | |
| NodeBinding(DummyChecker([PointerClass(float32_t)], [PointerClass(float32_t)]), | ||
| FloatGlobalMaxPoolTemplate.referenceTemplate, BasicTransformer) | ||
| ] | ||
|
|
||
| BasicCol2ImBindings = [ | ||
| NodeBinding( | ||
| PassThroughTypeChecker([PointerClass(type), PointerClass(int32_t), | ||
| PointerClass(int32_t)], [PointerClass(type)]), Col2ImTemplate.referenceTemplate, | ||
| BasicTransformer) for type in (int8_t, uint8_t, float32_t) | ||
| ] | ||
|
|
||
| BasicScatterBindings = [ | ||
| NodeBinding( | ||
| PassThroughTypeChecker( | ||
| [PointerClass(type), PointerClass(int32_t), PointerClass(type)], [PointerClass(type)]), | ||
| ScatterTemplate.referenceTemplate, BasicTransformer) for type in (int8_t, uint8_t, float32_t) | ||
| ] | ||
|
|
||
| BasicResizeBindings = [ | ||
| NodeBinding(PassThroughTypeChecker([PointerClass(type)], [PointerClass(type)]), ResizeTemplate.referenceTemplate, | ||
| BasicTransformer) for type in (int8_t, uint8_t, float32_t) | ||
| ] | ||
Uh oh!
There was an error while loading. Please reload this page.