@@ -1727,13 +1727,21 @@ class SparseShape {
17271727 const bool unit_external =
17281728 (tile_norms_.range ().rank () + 1u == nfused + gemm_helper.left_rank ());
17291729 const unsigned int u = unit_external ? 1u : 0u ;
1730+ // a fused broadcast on the RIGHT carries a SYNTHETIC unit right-external
1731+ // mode in the GemmHelper only (mirror of unit_external; see
1732+ // ContEngine::synthetic_unit_right_external); detect it from the one-rank
1733+ // mismatch with the actual right norm tensor and pad the folded
1734+ // right/result views with a unit extent
1735+ const bool unit_right_external = (other.tile_norms_ .range ().rank () + 1u ==
1736+ nfused + gemm_helper.right_rank ());
1737+ const unsigned int u_right = unit_right_external ? 1u : 0u ;
17301738
17311739 // check that the ranks match the folded gemm ranks plus the fused modes,
17321740 // and that the fused and contracted mode extents of the two shapes are
17331741 // congruent
17341742 TA_ASSERT (tile_norms_.range ().rank () + u ==
17351743 nfused + gemm_helper.left_rank ());
1736- TA_ASSERT (other.tile_norms_ .range ().rank () ==
1744+ TA_ASSERT (other.tile_norms_ .range ().rank () + u_right ==
17371745 nfused + gemm_helper.right_rank ());
17381746 for (unsigned int d = 0u ; d < nfused; ++d)
17391747 TA_ASSERT (left_extent[d] == right_extent[d]);
@@ -1751,14 +1759,16 @@ class SparseShape {
17511759 for (unsigned int i = gemm_helper.left_inner_begin ();
17521760 i < gemm_helper.left_inner_end (); ++i)
17531761 K *= left_extent[nfused + i - u];
1754- for (unsigned int i = gemm_helper.right_outer_begin ();
1755- i < gemm_helper.right_outer_end (); ++i)
1756- N *= right_extent[nfused + i];
1762+ if (!unit_right_external)
1763+ for (unsigned int i = gemm_helper.right_outer_begin ();
1764+ i < gemm_helper.right_outer_end (); ++i)
1765+ N *= right_extent[nfused + i];
17571766
17581767 // result size vectors: fused modes (from this), then the left and right
1759- // outer modes (the synthetic unit left-external mode is absent from the
1760- // actual result)
1761- const unsigned int result_rank = nfused + gemm_helper.result_rank () - u;
1768+ // outer modes (the synthetic unit left-/right-external modes are absent
1769+ // from the actual result)
1770+ const unsigned int result_rank =
1771+ nfused + gemm_helper.result_rank () - u - u_right;
17621772 std::shared_ptr<vector_type> result_size_vectors (
17631773 new vector_type[result_rank], std::default_delete<vector_type[]>());
17641774 unsigned int x = 0ul ;
@@ -1768,9 +1778,10 @@ class SparseShape {
17681778 for (unsigned int i = gemm_helper.left_outer_begin ();
17691779 i < gemm_helper.left_outer_end (); ++i, ++x)
17701780 result_size_vectors.get ()[x] = size_vectors_.get ()[nfused + i];
1771- for (unsigned int i = gemm_helper.right_outer_begin ();
1772- i < gemm_helper.right_outer_end (); ++i, ++x)
1773- result_size_vectors.get ()[x] = other.size_vectors_ .get ()[nfused + i];
1781+ if (!unit_right_external)
1782+ for (unsigned int i = gemm_helper.right_outer_begin ();
1783+ i < gemm_helper.right_outer_end (); ++i, ++x)
1784+ result_size_vectors.get ()[x] = other.size_vectors_ .get ()[nfused + i];
17741785
17751786 // the result norm tensor over (fused..., left outer..., right outer...)
17761787 using range_type = typename Tensor<value_type>::range_type;
@@ -1788,22 +1799,28 @@ class SparseShape {
17881799 lobounds.push_back (tile_norms_.range ().lobound_data ()[nfused + i]);
17891800 upbounds.push_back (tile_norms_.range ().upbound_data ()[nfused + i]);
17901801 }
1791- for (unsigned int i = gemm_helper.right_outer_begin ();
1792- i < gemm_helper.right_outer_end (); ++i) {
1793- lobounds.push_back (other.tile_norms_ .range ().lobound_data ()[nfused + i]);
1794- upbounds.push_back (other.tile_norms_ .range ().upbound_data ()[nfused + i]);
1795- }
1802+ if (!unit_right_external)
1803+ for (unsigned int i = gemm_helper.right_outer_begin ();
1804+ i < gemm_helper.right_outer_end (); ++i) {
1805+ lobounds.push_back (
1806+ other.tile_norms_ .range ().lobound_data ()[nfused + i]);
1807+ upbounds.push_back (
1808+ other.tile_norms_ .range ().upbound_data ()[nfused + i]);
1809+ }
17961810 Tensor<value_type> result_norms (range_type (lobounds, upbounds), 0 );
17971811
17981812 // the range spanned by modes [nfused, rank) of \p r, rebased to zero
17991813 // lobounds (scratch view for the slab-batched norm GEMM)
18001814 auto fold_range = [nfused](const range_type& r,
1801- const bool prepend_unit = false ) {
1815+ const bool prepend_unit = false ,
1816+ const bool append_unit = false ) {
18021817 const auto * extent = r.extent_data ();
18031818 container::svector<index1_type> extents;
1804- extents.reserve (r.rank () - nfused + (prepend_unit ? 1u : 0u ));
1819+ extents.reserve (r.rank () - nfused + (prepend_unit ? 1u : 0u ) +
1820+ (append_unit ? 1u : 0u ));
18051821 if (prepend_unit) extents.push_back (1 );
18061822 extents.insert (extents.end (), extent + nfused, extent + r.rank ());
1823+ if (append_unit) extents.push_back (1 );
18071824 return range_type (extents);
18081825 };
18091826
@@ -1847,9 +1864,11 @@ class SparseShape {
18471864 // buffer, so the accumulation lands in place
18481865 auto left_folded =
18491866 left.reshape (fold_range (left.range (), unit_external), H);
1850- auto right_folded = right.reshape (fold_range (right.range ()), H);
1867+ auto right_folded = right.reshape (
1868+ fold_range (right.range (), false , unit_right_external), H);
18511869 auto result_folded = result_norms.reshape (
1852- fold_range (result_norms.range (), unit_external), H);
1870+ fold_range (result_norms.range (), unit_external, unit_right_external),
1871+ H);
18531872 result_folded.gemm (left_folded, right_folded, abs_factor, gemm_helper);
18541873
18551874 // Hard zero tiles that are below the zero threshold.
0 commit comments