Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 78 additions & 46 deletions ops/matmul-inl.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<TC>()) {
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);
}
}
});
}
Expand All @@ -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<TC>()) {
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);
}
}
});
}
Expand All @@ -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<TC>()) {
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);
}
}
});
}
Expand All @@ -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<TC>()) {
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);
}
}
});
}
Expand All @@ -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<TC>()) {
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);
}
}
});
}
Expand All @@ -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<TC>()) {
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);
}
}
});
}
Expand Down
6 changes: 4 additions & 2 deletions ops/matmul.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 <class Func>
Expand Down Expand Up @@ -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);
});
}

Expand Down
14 changes: 11 additions & 3 deletions util/threading_context.h
Original file line number Diff line number Diff line change
Expand Up @@ -212,12 +212,20 @@ template <class Func>
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); });
Expand Down
Loading