Skip to content
2 changes: 2 additions & 0 deletions backends/samsung/_passes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from .decompose_einsum import DecomposeEinsum
from .decompose_glu import DecomposeGlu
from .decompose_linalg_vector_norm import DecomposeLinalgVectorNorm
from .decompose_remainder import DecomposeRemainder
from .decompose_roll import DecomposeRoll
from .fold_qdq import FoldQDQPass
from .fuse_activation import FuseActivationPass
Expand All @@ -31,6 +32,7 @@
"DecomposeEinsum",
"DecomposeGlu",
"DecomposeLinalgVectorNorm",
"DecomposeRemainder",
"DecomposeRoll",
"FoldQDQPass",
"FuseActivationPass",
Expand Down
180 changes: 149 additions & 31 deletions backends/samsung/_passes/annotate_qparams.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,8 @@ class AnnotateQparamsPass(ExportPass):
exir_ops.edge.aten.cat.default,
exir_ops.edge.aten.expand_copy.default,
exir_ops.edge.aten.split_with_sizes_copy.default,
exir_ops.edge.aten.clone.default,
exir_ops.edge.aten.contiguous.default,
}

def __init__(self, edge_program: ExportedProgram):
Expand Down Expand Up @@ -84,40 +86,154 @@ def _impl(node: Node, res_list: List[Node]):
_impl(user, res_list)
return res_list

def _walk_qdq_chain_to_terminals(self, cur: Node) -> List[Node]:

@psiddh psiddh Oct 2, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we add focused graph-level tests for _walk_qdq_chain_to_terminals and the backward qparam propagation path? In particular: silu -> split/chunk, Q/DQ fanout, clone/contiguous, and a negative case where multi-input ops are not backward-propagated. If not in this PR, please track it as follow-up since any issues here could become silent accuracy regressions.

r"""Walk forward from a Q/DQ node `cur` through the Q-DQ chain,
returning every terminal node (last Q/DQ node before a non-Q/DQ
consumer). Handles fan-out: the SAME quantized tensor is commonly
dequantized into multiple branches (e.g. one Q feeding two DQs for
two independent consumers) -- each branch is walked and its own
terminal is collected, instead of the chain silently stopping at
`cur` when it has more than one Q/DQ child.

Mirrors the DFS shape of `_get_last_dqs`, applied starting from a
single Q/DQ node rather than its non-Q/DQ source.
"""
next_nodes = [
u
for u in cur.users
if u.target in QuantConstants.QUANT_OPS_KEY_MAP
or u.target in QuantConstants.DEQUANT_OPS_KEY_MAP
]
terminals: List[Node] = []
""" `cur` may feed a non-Q/DQ consumer *and* further Q/DQ branches. It
still terminates the chain for that direct consumer, so collect it
as well as walking the branches."""
if not next_nodes or len(next_nodes) < len(cur.users):
terminals.append(cur)
for nxt in next_nodes:
terminals.extend(self._walk_qdq_chain_to_terminals(nxt))
return terminals

def _collect_dq_nodes(self, node: Node) -> List[Node]:
"""For each user of `node`, resolve it to a propagate candidate: a
non-Q user is used directly, a Q user is walked through its Q-DQ
chain (including fan-out) to find every terminal node."""
dq_nodes: List[Node] = []
for user in node.users:
if user.target not in QuantConstants.QUANT_OPS_KEY_MAP:
# If user is a direct propagate node (not Q), collect it too
dq_nodes.append(user)
continue
# user is a Q node: walk through the Q-DQ chain (including any
# fan-out branches) to find every terminal node.
dq_nodes.extend(self._walk_qdq_chain_to_terminals(user))
return dq_nodes

def _is_propagatable(self, candidate: Node) -> bool:
"""True if `candidate` is a SharedQuant propagate node that is safe
to annotate with quantize_attrs and recurse into: it must be a
propagate node, and if it has exactly one user, that user must not
be a Q/DQ (a Q/DQ boundary already carries its own quant params)."""
if candidate.target not in self.propagate_nodes:
return False
if len(candidate.users) == 1:
only_user = next(iter(candidate.users))
if (
only_user.target in QuantConstants.QUANT_OPS_KEY_MAP
or only_user.target in QuantConstants.DEQUANT_OPS_KEY_MAP
):
return False
return True

def _propagate_into_dequant_users(self, dq_node: Node, user_attrs) -> None:
"""dq_node is a DQ: propagate to each of its users that is itself an
eligible (non-Q/DQ) propagate node."""
for op_user in dq_node.users:
if (
op_user.target in QuantConstants.QUANT_OPS_KEY_MAP
or op_user.target in QuantConstants.DEQUANT_OPS_KEY_MAP
):
continue
if not self._is_propagatable(op_user):
continue
op_user.meta["quantize_attrs"] = user_attrs
self._propagate_quant_params(op_user)

def _propagate_quant_params(self, node: Node):
assert (
quantize_attrs := node.meta.get("quantize_attrs")
), "Must be annotated node."
requantize_map: Dict[Node, Node] = node.meta.get("requantize", {})
while node.users:
if len(node.users) != 1:
break
user = list(node.users.keys())[0]
if (
user.target not in QuantConstants.QUANT_OPS_KEY_MAP
and user.target not in QuantConstants.DEQUANT_OPS_KEY_MAP
):
break
node = user
# Case1: ...-q-dq(cur)-propagate_node-node(not d-dq)
# Case2: propagate_node(propagateed)-propagate_node-node(not q-dq)
for idx, user in enumerate(node.users.keys()):
# Walk through Q-DQ chains, handling multiple Q-DQ branches.
# For node->Q->DQ->op1 and node->Q->DQ->op3, we collect all last DQ nodes.
dq_nodes = self._collect_dq_nodes(node)
# Case1: ...-q-dq(cur)-propagate_node-node(not q-dq)
# Case2: propagate_node(propagated)-propagate_node-node(not q-dq)
for idx, dq_node in enumerate(dq_nodes):
# For the branch who need to be requantized, we propagate the requantize params
user_attrs = requantize_map.get(idx, quantize_attrs)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

requantize_map is keyed by a different list than the one being indexed here.

_annotate_requantize builds requantize_map keyed by idx into self._get_last_dqs(node) (line 256), which yields only Q/DQ-chain terminals. This lookup indexes by position in _collect_dq_nodes(node), which additionally includes every non-Q user verbatim (lines 120-122).

The two lists agree only when the source node has no non-Q users. Add a direct float consumer alongside the quantized branches and they shift relative to each other, so a requantized branch picks up another branch's params.

Keying requantize by node identity instead of positional index would make this robust; at minimum both call sites should walk the same list.


Generated by an AI reviewer (Claude Code). Please verify before acting on it.

if user.target not in self.propagate_nodes:
if dq_node.target in QuantConstants.DEQUANT_OPS_KEY_MAP:
self._propagate_into_dequant_users(dq_node, user_attrs)
elif self._is_propagatable(dq_node):
# dq_node is not a DQ but a propagate node directly connected to source
dq_node.meta["quantize_attrs"] = user_attrs
self._propagate_quant_params(dq_node)

def _backward_propagate(self, node: Node):
"""Walk backward from `node`, copying its quantize_attrs into unannotated
single-input upstream ops.

Handles patterns like `DQ → SiLU → chunk → Q` where forward propagation
cannot populate SiLU's quantize_attrs because SiLU is not in
`propagate_nodes` (its input/output scales differ in general). But when
the downstream is a SharedQuant op (chunk/split/view/permute/...) whose
input scale must equal its output scale, the intermediate op's output
scale is fully determined by the downstream shared scale, so backward
propagation is safe.

Stops at:
- already-annotated upstream (respect forward pass results)
- Q/DQ boundaries (scale is defined by the Q/DQ params themselves)
- non-call_function nodes (placeholders, get_attr, output)
- multi-input upstream ops (ambiguous which input to follow)
"""
quant_attrs = node.meta.get("quantize_attrs")
if not quant_attrs:
return
inputs = node.all_input_nodes
if len(inputs) != 1:
return
upstream = inputs[0]
if upstream.meta.get("quantize_attrs"):
return
if upstream.target in QuantConstants.QUANT_OPS_KEY_MAP:
return
if upstream.target in QuantConstants.DEQUANT_OPS_KEY_MAP:
return
if upstream.op != "call_function":
return
upstream.meta["quantize_attrs"] = quant_attrs
self._backward_propagate(upstream)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The backward walk recurses one hop too far.

upstream.meta["quantize_attrs"] = quant_attrs
self._backward_propagate(upstream)

The docstring's justification holds for the first hop: for SiLU → chunk, chunk's shared scale is also the scale of chunk's input, i.e. SiLU's output.

The recursion reapplies that reasoning one level up, where the premise no longer holds. For conv → SiLU → chunk, conv's output is SiLU's input, whose scale is unrelated to SiLU's output scale — SiLU is excluded from propagate_nodes for exactly that reason. Annotating conv with chunk's scale is wrong.

The if upstream.meta.get("quantize_attrs"): return guard usually saves this in a fully quantized graph, since conv would already be annotated by the forward pass. But whenever two consecutive non-identity ops are unannotated, this propagates a scale that doesn't belong to them.

Suggest dropping the recursive call, or continuing only while upstream is itself in propagate_nodes.


Generated by an AI reviewer (Claude Code). Please verify before acting on it.


def _propagate_quant_params_backward_all(self, graph_module: GraphModule):
"""For every SharedQuant propagate node with annotated quantize_attrs,
walk backward and fill in unannotated single-input upstream ops.

Multi-input propagate ops (cat/concat) are excluded because their
upstream is ambiguous — different input branches may legitimately have
different scales, and picking one to backward-propagate would corrupt
the others.
"""
single_input_shared = self.propagate_nodes - {
exir_ops.edge.aten.concat.default,
exir_ops.edge.aten.cat.default,
}
for node in graph_module.graph.nodes:
if node.target not in single_input_shared:
continue
if len(user.users) == 1:
# Possibily no need for checking len(users)>1
user_of_user = list(user.users)[0]
# node-q-dq-propagate-q-dq not need for propagatey
if (
user_of_user.target in QuantConstants.QUANT_OPS_KEY_MAP
or user_of_user.target in QuantConstants.DEQUANT_OPS_KEY_MAP
):
continue
# propagate quant for node-q-dq-propagate_node-node(not qdq)
user.meta["quantize_attrs"] = user_attrs
self._propagate_quant_params(user)
if not node.meta.get("quantize_attrs"):
continue
self._backward_propagate(node)

def _annotate_requantize(self, node: Node):
assert (
Expand Down Expand Up @@ -185,12 +301,13 @@ def _annotate(self, graph_module: GraphModule):
):
# Currently, don't add quant info for d_qd node here.
continue
elif source_node.target == operator.getitem:
source_node.meta["quantize_attrs"] = quant_attrs
source_node = source_node.args[0]

source_node.meta["quantize_attrs"] = quant_attrs
self._annotate_requantize(source_node)
if source_node.target == operator.getitem:
source_node.args[0].meta["quantize_attrs"] = quant_attrs
self._annotate_requantize(source_node.args[0])
else:
self._annotate_requantize(source_node)

self._propagate_quant_params(source_node)

def _annotate_in_quantize_attrs(self, graph_module: GraphModule):
Expand Down Expand Up @@ -244,6 +361,7 @@ def _annotate_decomposed_mm(self, graph_module: GraphModule):

def call(self, graph_module: GraphModule):
self._annotate(graph_module)
self._propagate_quant_params_backward_all(graph_module)
self._annotate_decomposed_mm(graph_module)
self._annotate_in_quantize_attrs(graph_module)
graph_module.recompile()
Expand Down
126 changes: 126 additions & 0 deletions backends/samsung/_passes/decompose_remainder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
# Copyright (c) Qualcomm Innovation Center, Inc.
# Copyright (c) Samsung Electronics Co. LTD
# All rights reserved
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

import torch
from executorch.exir.dialects._ops import ops as exir_ops
from executorch.exir.dialects.edge._ops import EdgeOpOverload
from executorch.exir.pass_base import ExportPass, PassResult
from torchao.quantization.pt2e.utils import get_new_attr_name_with_prefix

from .utils import copy_meta, create_const_node


class DecomposeRemainder(ExportPass):
"""
Decompose remainder.Scalar and remainder.Tensor using the identity:
remainder(x, y) = x - floor(x / y) * y
"""

def __init__(self):
super(DecomposeRemainder, self).__init__()
self.remainder_targets = {
torch.ops.aten.remainder.Scalar,
torch.ops.aten.remainder.Scalar_Tensor,
torch.ops.aten.remainder.Tensor,
exir_ops.edge.aten.remainder.Scalar,
exir_ops.edge.aten.remainder.Scalar_Tensor,
exir_ops.edge.aten.remainder.Tensor,
}

def call(self, graph_module: torch.fx.GraphModule):
graph = graph_module.graph
# Cache scalar:node mappings to avoid duplicate buffer registrations if the same scalar divisor appears in multiple remainder ops
const_cache = {}
for node in list(graph.nodes):
if node.op == "call_function" and node.target in self.remainder_targets:
x_arg = node.args[0]
y_arg = node.args[1]
is_edge = isinstance(node.target, EdgeOpOverload)
meta = node.meta

val = meta.get("val", None)
if val is None or val.dtype.is_floating_point:
continue

floor_div_op = (
exir_ops.edge.aten.floor_divide.default
if is_edge
else torch.ops.aten.floor_divide.default
)
mul_op = (
exir_ops.edge.aten.mul.Tensor
if is_edge
else torch.ops.aten.mul.Tensor
)
sub_op = (
exir_ops.edge.aten.sub.Tensor
if is_edge
else torch.ops.aten.sub.Tensor
Comment thread
Jiseong-oh marked this conversation as resolved.
)

is_x_scalar = not isinstance(x_arg, torch.fx.Node)
if is_x_scalar and is_edge:
if x_arg not in const_cache:
attr_name = get_new_attr_name_with_prefix("_remainder_const_")(
graph_module
)
const_cache[x_arg] = create_const_node(
graph, graph_module, attr_name, x_arg, node
)
x_node = const_cache[x_arg]
else:
x_node = x_arg

is_y_scalar = not isinstance(y_arg, torch.fx.Node)
if is_y_scalar and is_edge:
if y_arg not in const_cache:
attr_name = get_new_attr_name_with_prefix("_remainder_const_")(
graph_module
)
const_cache[y_arg] = create_const_node(
graph, graph_module, attr_name, y_arg, node
)
y_node = const_cache[y_arg]
else:
y_node = y_arg

with graph.inserting_before(node):
floor_div_node = graph.create_node(
"call_function", floor_div_op, (x_node, y_node)
)
Comment thread
Jiseong-oh marked this conversation as resolved.
floor_div_node.meta = copy_meta(meta)

mul_node = graph.create_node(
"call_function", mul_op, (floor_div_node, y_node)
)
mul_node.meta = copy_meta(meta)

sub_node = graph.create_node(
"call_function", sub_op, (x_node, mul_node)
)
sub_node.meta = copy_meta(meta)

# Cast back to the original integer dtype, which the
# division may have promoted away from.
to_copy_op = (
exir_ops.edge.aten._to_copy.default
if is_edge
else torch.ops.aten._to_copy.default
)
cast_node = graph.create_node(
"call_function",
to_copy_op,
(sub_node,),
{"dtype": val.dtype},
)
cast_node.meta = copy_meta(meta)
for user in node.users.copy():
user.replace_input_with(node, cast_node)

graph.eliminate_dead_code()
graph_module.recompile()
return PassResult(graph_module, True)
2 changes: 2 additions & 0 deletions backends/samsung/_passes/enn_pass_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
DecomposeEinsum,
DecomposeGlu,
DecomposeLinalgVectorNorm,
DecomposeRemainder,
DecomposeRoll,
FoldQDQPass,
FuseActivationPass,
Expand Down Expand Up @@ -59,6 +60,7 @@ def transform_for_annotation_pass(self, graph_module: GraphModule):
def transform_for_export_pass(self, exported_program: ExportedProgram):
self.add_pass(ComputeConstAttrs())
self.add_pass(DecomposeRoll())
self.add_pass(DecomposeRemainder())
self._transform(exported_program.graph_module)
return exported_program

Expand Down
Loading
Loading