From 298dabcc179831c3060a017e676886e81ad724ae Mon Sep 17 00:00:00 2001 From: "The gemma.cpp Authors" Date: Sun, 19 Jul 2026 09:35:52 -0700 Subject: [PATCH] Fix crash/abort in matmul when partition size exceeds kMaxNC. PiperOrigin-RevId: 950420776 --- ops/matmul-inl.h | 124 ++++++++++++++++++++++++--------------- ops/matmul.h | 6 +- util/threading_context.h | 14 ++++- 3 files changed, 93 insertions(+), 51 deletions(-) diff --git a/ops/matmul-inl.h b/ops/matmul-inl.h index d6c4d8f8..7ccb4d3d 100644 --- a/ops/matmul-inl.h +++ b/ops/matmul-inl.h @@ -924,16 +924,21 @@ class MMLoops { MMKernel::B3A2C0(A, B, range_mc, range_kc, range_nc, args, MMSetC(), C.View(0, range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, - MMSetC(), C2); - } - if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } } }); } @@ -958,17 +963,22 @@ class MMLoops { A, B, range_mc, args.ranges_kc, range_nc, args, C.View(0, range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, - range_nc, args, C2); - } - if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, + range_nc, args, C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, + range_nc, args, C2); + } } }); } @@ -994,15 +1004,21 @@ class MMLoops { A, B, range_mc, range_kc, range_nc, args, MMSetC(), C.View(range_mc.begin(), range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, - MMSetC(), C2); - } if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } } }); } @@ -1026,16 +1042,21 @@ class MMLoops { A, B, range_mc, args.ranges_kc, range_nc, args, C.View(range_mc.begin(), range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, - args, C2); - } - if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, + args, C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, + args, C2); + } } }); } @@ -1060,15 +1081,21 @@ class MMLoops { A, B, range_mc, range_kc, range_nc, args, MMSetC(), C.View(range_mc.begin(), range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, - MMSetC(), C2); - } if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::B3A2C0(A, *B2, range_mc, range_kc, range_nc, args, + MMSetC(), C2); + } } }); } @@ -1091,16 +1118,21 @@ class MMLoops { A, B, range_mc, args.ranges_kc, range_nc, args, C.View(range_mc.begin(), range_nc.begin(), range_nc.Num())); - const StridedViewBF C2 = args.env.C_tiles.C( - Extents2D(range_mc.Num(), range_nc.Num()), worker); - - if (B2 != nullptr) { - MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, - args, C2); - } - if constexpr (IsBF16()) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + if (B2 != nullptr) { + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, + args, C2); + } args.options.MaybeCallFunc(C, range_mc, range_nc, C2, worker); + } else { + if (B2 != nullptr) { + const StridedViewBF C2 = args.env.C_tiles.C( + Extents2D(range_mc.Num(), range_nc.Num()), worker); + MMKernel::ForeachKC(A, *B2, range_mc, args.ranges_kc, range_nc, + args, C2); + } } }); } diff --git a/ops/matmul.h b/ops/matmul.h index 9e0d3f5d..3071cec5 100644 --- a/ops/matmul.h +++ b/ops/matmul.h @@ -131,7 +131,8 @@ struct MMParallelWithinCluster { range_n, n_multiple, inner_tasks, ctx, cluster_idx, caller, [&](const IndexRange& worker_range, size_t worker) { func(worker_range, worker); - }); + }, + kMaxNC); } template @@ -212,7 +213,8 @@ struct MMParallelHierarchical { cluster_range, n_multiple, inner_tasks, ctx, cluster_idx, caller, [&](const IndexRange& worker_range, size_t worker) { func(worker_range, worker); - }); + }, + kMaxNC); }); } diff --git a/util/threading_context.h b/util/threading_context.h index 7e595ba6..5f4f0df2 100644 --- a/util/threading_context.h +++ b/util/threading_context.h @@ -212,12 +212,20 @@ template void ParallelPartitionWithinCluster(const IndexRange range, size_t task_multiple, size_t inner_tasks, ThreadingContext& ctx, size_t cluster_idx, - hwy::pool::Caller caller, - const Func& func) { + hwy::pool::Caller caller, const Func& func, + size_t max_size = 0) { HWY_DASSERT(1 <= inner_tasks && inner_tasks <= 4); const size_t num_workers = ctx.pools.Cluster(cluster_idx).NumWorkers(); + size_t tasks = num_workers * inner_tasks; + if (max_size != 0) { + const size_t target_size = hwy::RoundDownTo(max_size, task_multiple); + if (target_size != 0) { + const size_t min_tasks = hwy::DivCeil(range.Num(), target_size); + tasks = std::max(tasks, min_tasks); + } + } const IndexRangePartition ranges = - StaticPartition(range, num_workers * inner_tasks, task_multiple); + StaticPartition(range, tasks, task_multiple); ParallelForWithinCluster( ranges.NumTasks(), ctx, cluster_idx, caller, [&](uint64_t task, size_t worker) { func(ranges.Range(task), worker); });