Skip to content
Merged
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
4 changes: 3 additions & 1 deletion CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
cmake_minimum_required(VERSION 3.16)
project(CSR4MPI LANGUAGES CXX)

set(CSR4MPI_VALUE_TYPE "1" CACHE STRING "Scalar type: 0=float,1=double,2=complex<float>,3=complex<double>")
# CSR4MPI_VALUE_TYPE is now deprecated - all types are compiled into the library.
# This option is kept for backward compatibility but has no effect.
set(CSR4MPI_VALUE_TYPE "1" CACHE STRING "DEPRECATED: All scalar types are now always compiled. This option has no effect.")
option(CSR4MPI_ENABLE_BLAS "Enable optional BLAS placeholder path" OFF)
option(CSR4MPI_ENABLE_OPENMP "Enable OpenMP parallel kernels" ON)

Expand Down
19 changes: 10 additions & 9 deletions bench/bench_large_spmv.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,9 @@
#include <vector>

using namespace csr4mpi;
using Scalar = double;

static cCSRMatrix ExtractLocal(const cCSRMatrix& full, const cRowDistribution& dist)
static cCSRMatrix<Scalar> ExtractLocal(const cCSRMatrix<Scalar>& full, const cRowDistribution& dist)
{
iIndex rBeg = dist.iGlobalRowBegin();
iIndex rEnd = dist.iGlobalRowEnd();
Expand All @@ -20,7 +21,7 @@ static cCSRMatrix ExtractLocal(const cCSRMatrix& full, const cRowDistribution& d
const auto& FV = full.vValues();
std::vector<iIndex> lRP(localRows + 1, 0);
std::vector<iIndex> lC;
std::vector<vScalar> lV;
std::vector<Scalar> lV;
for (iSize lr = 0; lr < localRows; ++lr) {
iIndex g = rBeg + lr;
for (iIndex k = FRP[(size_t)g]; k < FRP[(size_t)g + 1]; ++k) {
Expand All @@ -29,7 +30,7 @@ static cCSRMatrix ExtractLocal(const cCSRMatrix& full, const cRowDistribution& d
}
lRP[(size_t)lr + 1] = (iIndex)lC.size();
}
cCSRMatrix local(rBeg, rEnd, full.iGlobalColCount(), lRP, lC, lV, full.eSymmetry());
cCSRMatrix<Scalar> local(rBeg, rEnd, full.iGlobalColCount(), lRP, lC, lV, full.eSymmetry());
local.AttachDistribution(std::make_shared<cRowDistribution>(dist));
return local;
}
Expand All @@ -43,25 +44,25 @@ int main(int argc, char** argv)
std::string matrixFile = (argc > 1 ? argv[1] : "bcsstk13.mtx");
std::string path = std::string(CSR4MPI_SOURCE_DIR) + "/tests/data/" + matrixFile;
std::vector<iIndex> gRP, gCI;
std::vector<vScalar> gV;
std::vector<Scalar> gV;
iIndex gRows = 0, gCols = 0;
if (!LoadMatrixMarket(path, gRP, gCI, gV, gRows, gCols)) {
if (!LoadMatrixMarket<Scalar>(path, gRP, gCI, gV, gRows, gCols)) {
if (rank == 0)
std::cout << "Matrix not found: " << path << "\n";
MPI_Finalize();
return 0;
}
auto dist = cRowDistribution::CreateBlockDistribution(gRows, worldSize, rank);
cCSRMatrix local = ExtractLocal(cCSRMatrix(0, gRows, gCols, gRP, gCI, gV), dist);
cCSRMatrix<Scalar> local = ExtractLocal(cCSRMatrix<Scalar>(0, gRows, gCols, gRP, gCI, gV), dist);
// Build local slice of x (random deterministic)
const auto& offs = dist.vRowOffsets();
iIndex cBeg = offs[(size_t)rank];
iIndex cEnd = offs[(size_t)rank + 1];
std::vector<vScalar> xLocal;
std::vector<Scalar> xLocal;
xLocal.reserve(static_cast<size_t>(cEnd - cBeg));
for (iIndex g = cBeg; g < cEnd; ++g)
xLocal.push_back(static_cast<vScalar>(1));
std::vector<vScalar> yLocal;
xLocal.push_back(static_cast<Scalar>(1));
std::vector<Scalar> yLocal;
int iters = (argc > 2 ? std::stoi(argv[2]) : 10);
MPI_Barrier(MPI_COMM_WORLD);
double t0 = MPI_Wtime();
Expand Down
21 changes: 11 additions & 10 deletions bench/bench_operations.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,16 +6,17 @@
#include <vector>

using namespace csr4mpi;
using Scalar = double;

static cCSRMatrix BuildRandomUniform(iSize rows, iSize cols, iSize nnzPerRow, unsigned seed)
static cCSRMatrix<Scalar> BuildRandomUniform(iSize rows, iSize cols, iSize nnzPerRow, unsigned seed)
{
std::mt19937_64 gen(seed);
std::uniform_int_distribution<iSize> colDist(0, cols - 1);
std::uniform_real_distribution<double> valDist(0.0, 1.0);
std::vector<iIndex> vRowPtr(rows + 1, 0);
std::vector<iIndex> vColInd;
vColInd.reserve(rows * nnzPerRow);
std::vector<vScalar> vValues;
std::vector<Scalar> vValues;
vValues.reserve(rows * nnzPerRow);
for (iSize r = 0; r < rows; ++r) {
std::vector<iSize> colsChosen;
Expand All @@ -36,11 +37,11 @@ static cCSRMatrix BuildRandomUniform(iSize rows, iSize cols, iSize nnzPerRow, un
for (auto c : colsChosen) {
vColInd.push_back(static_cast<iIndex>(c));
double rv = valDist(gen);
vValues.push_back(static_cast<vScalar>(rv));
vValues.push_back(static_cast<Scalar>(rv));
}
vRowPtr[static_cast<size_t>(r + 1)] = static_cast<iIndex>(vColInd.size());
}
return cCSRMatrix(0, rows, cols, vRowPtr, vColInd, vValues);
return cCSRMatrix<Scalar>(0, rows, cols, vRowPtr, vColInd, vValues);
}

int main(int argc, char** argv)
Expand All @@ -62,12 +63,12 @@ int main(int argc, char** argv)
repeats = std::stoi(argv[5]);

auto A = BuildRandomUniform(rows, cols, nnzPerRow, 42);
std::vector<vScalar> x(static_cast<size_t>(cols));
std::vector<Scalar> x(static_cast<size_t>(cols));
for (size_t i = 0; i < x.size(); ++i)
x[i] = static_cast<vScalar>(1);
x[i] = static_cast<Scalar>(1);

// SpMV benchmark
std::vector<vScalar> y;
std::vector<Scalar> y;
auto t0 = std::chrono::high_resolution_clock::now();
for (int r = 0; r < repeats; ++r) {
SpMV(A, x, y);
Expand All @@ -76,10 +77,10 @@ int main(int argc, char** argv)
std::chrono::duration<double> dtSpMV = t1 - t0;

// SpMM benchmark
std::vector<vScalar> X(static_cast<size_t>(cols * spmmCols));
std::vector<Scalar> X(static_cast<size_t>(cols * spmmCols));
for (size_t i = 0; i < X.size(); ++i)
X[i] = static_cast<vScalar>(1);
std::vector<vScalar> Y;
X[i] = static_cast<Scalar>(1);
std::vector<Scalar> Y;
auto t2 = std::chrono::high_resolution_clock::now();
for (int r = 0; r < repeats; ++r) {
SpMM(A, X, spmmCols, Y);
Expand Down
26 changes: 24 additions & 2 deletions include/BlasAdapter.h
Original file line number Diff line number Diff line change
@@ -1,9 +1,31 @@
#pragma once
#include "CSRMatrix.h"
#include "Operations.h"
#include "Global.h"
#include <vector>
#include <iostream>

namespace csr4mpi {
bool bBlasEnabled();
void SpMMBlas(const cCSRMatrix& A, const std::vector<vScalar>& X, iSize nCols, std::vector<vScalar>& Y);

inline bool bBlasEnabled()
{
#ifdef CSR4MPI_USE_BLAS
return true;
#else
return false;
#endif
}

template <typename Scalar>
void SpMMBlas(const cCSRMatrix<Scalar>& A, const std::vector<Scalar>& X, iSize nCols, std::vector<Scalar>& Y)
{
static_assert(is_supported_scalar_v<Scalar>, "Scalar must be float, double, std::complex<float>, or std::complex<double>");
#ifdef CSR4MPI_USE_BLAS
// Placeholder: in future, detect dense conversion threshold and call cblas_Xgemm.
// For now, just log once per call and fallback.
std::cerr << "[BLAS placeholder] Falling back to internal SpMM path.\n";
#endif
SpMM(A, X, nCols, Y);
}

}
157 changes: 151 additions & 6 deletions include/CSRComm.h
Original file line number Diff line number Diff line change
@@ -1,18 +1,163 @@
#pragma once

#include "Global.h"
#include "CSRMatrix.h"
#include "CommPattern.h"
#include <mpi.h>
#include <vector>

namespace csr4mpi {
class cCSRMatrix;
class cCommPattern;

template <typename Scalar>
class cCSRComm {
static_assert(is_supported_scalar_v<Scalar>, "Scalar must be float, double, std::complex<float>, or std::complex<double>");

public:
static void Assemble(cCSRMatrix& cLocal,
const cCommPattern& cPattern,
const std::vector<vScalar>& vValues,
MPI_Comm cComm);
static void Assemble(cCSRMatrix<Scalar>& cLocal,
const cCommPattern<Scalar>& cPattern,
const std::vector<Scalar>& vValues,
MPI_Comm cComm)
{
const std::vector<iIndex>& vSendRows = cPattern.vSendRows();
const std::vector<iIndex>& vSendCols = cPattern.vSendCols();
const std::vector<iSize>& vSendOffsets = cPattern.vSendOffsets();
const std::vector<int>& vTargets = cPattern.vTargetRanks();

const std::vector<iIndex>& vRowPtr = cLocal.vRowPtr();
const std::vector<iIndex>& vColInd = cLocal.vColInd();
std::vector<Scalar>& vLocalVal = cLocal.vValues();

(void)vSendOffsets;
(void)vTargets;

int iRank = 0;
int iWorldSize = 1;
MPI_Comm_rank(cComm, &iRank);
MPI_Comm_size(cComm, &iWorldSize);

// Prepare send counts per target.
std::vector<int> vSendCounts(iWorldSize, 0);
std::vector<int> vRecvCounts(iWorldSize, 0);

for (std::size_t i = 0; i < vTargets.size(); ++i) {
int iTarget = vTargets[i];
iSize iBegin = vSendOffsets[i];
iSize iEnd = vSendOffsets[i + 1];
vSendCounts[static_cast<std::size_t>(iTarget)] += static_cast<int>(iEnd - iBegin);
}

// Exchange counts.
MPI_Alltoall(vSendCounts.data(), 1, MPI_INT,
vRecvCounts.data(), 1, MPI_INT,
cComm);

int iTotalSend = 0;
int iTotalRecv = 0;
for (int i = 0; i < iWorldSize; ++i) {
iTotalSend += vSendCounts[static_cast<std::size_t>(i)];
iTotalRecv += vRecvCounts[static_cast<std::size_t>(i)];
}

struct cTriplet {
iIndex m_iRow;
iIndex m_iCol;
Scalar m_vVal;
};

std::vector<cTriplet> vSendBuf(static_cast<std::size_t>(iTotalSend));
std::vector<cTriplet> vRecvBuf(static_cast<std::size_t>(iTotalRecv));

// Fill send buffer in rank-major order.
std::vector<int> vOffsets(iWorldSize, 0);
for (int i = 1; i < iWorldSize; ++i) {
vOffsets[static_cast<std::size_t>(i)] = vOffsets[static_cast<std::size_t>(i - 1)] + vSendCounts[static_cast<std::size_t>(i - 1)];
}

for (std::size_t iBucket = 0; iBucket < vTargets.size(); ++iBucket) {
int iTarget = vTargets[iBucket];
iSize iBegin = vSendOffsets[iBucket];
iSize iEnd = vSendOffsets[iBucket + 1];

for (iSize i = iBegin; i < iEnd; ++i) {
int iPos = vOffsets[static_cast<std::size_t>(iTarget)]++;
vSendBuf[static_cast<std::size_t>(iPos)].m_iRow = vSendRows[static_cast<std::size_t>(i)];
vSendBuf[static_cast<std::size_t>(iPos)].m_iCol = vSendCols[static_cast<std::size_t>(i)];
vSendBuf[static_cast<std::size_t>(iPos)].m_vVal = vValues[static_cast<std::size_t>(i)];
}
}

// Build displacement arrays for Alltoallv.
std::vector<int> vSendDispls(iWorldSize, 0);
std::vector<int> vRecvDispls(iWorldSize, 0);

for (int i = 1; i < iWorldSize; ++i) {
vSendDispls[static_cast<std::size_t>(i)] = vSendDispls[static_cast<std::size_t>(i - 1)] + vSendCounts[static_cast<std::size_t>(i - 1)];
vRecvDispls[static_cast<std::size_t>(i)] = vRecvDispls[static_cast<std::size_t>(i - 1)] + vRecvCounts[static_cast<std::size_t>(i - 1)];
}

// Define MPI datatype for cTriplet with correct scalar handling.
MPI_Datatype tTripletType;
MPI_Datatype tCreatedScalarType = MPI_DATATYPE_NULL;
MPI_Datatype tScalarType = mpi_helper::GetMPIDatatype<Scalar>(&tCreatedScalarType);

{
cTriplet cDummy;
int iBlockLengths[3] = { 1, 1, 1 };
MPI_Aint vDispls[3];
MPI_Aint iBase;

MPI_Get_address(&cDummy, &iBase);
MPI_Get_address(&cDummy.m_iRow, &vDispls[0]);
MPI_Get_address(&cDummy.m_iCol, &vDispls[1]);
MPI_Get_address(&cDummy.m_vVal, &vDispls[2]);

vDispls[0] -= iBase;
vDispls[1] -= iBase;
vDispls[2] -= iBase;

MPI_Datatype vTypes[3] = { MPI_LONG_LONG, MPI_LONG_LONG, tScalarType };
MPI_Type_create_struct(3, iBlockLengths, vDispls, vTypes, &tTripletType);
MPI_Type_commit(&tTripletType);
}

MPI_Alltoallv(vSendBuf.data(),
vSendCounts.data(),
vSendDispls.data(),
tTripletType,
vRecvBuf.data(),
vRecvCounts.data(),
vRecvDispls.data(),
tTripletType,
cComm);

MPI_Type_free(&tTripletType);
if (tCreatedScalarType != MPI_DATATYPE_NULL) {
MPI_Type_free(&tCreatedScalarType);
}

// Accumulate received contributions into local CSR matrix.
for (const cTriplet& cT : vRecvBuf) {
iIndex iGlobalRow = cT.m_iRow;
iIndex iGlobalCol = cT.m_iCol;

iIndex iLocalRow = static_cast<iIndex>(iGlobalRow - cLocal.iGlobalRowBegin());
iIndex iStart = vRowPtr[static_cast<std::size_t>(iLocalRow)];
iIndex iEnd = vRowPtr[static_cast<std::size_t>(iLocalRow + 1)];

for (iIndex i = iStart; i < iEnd; ++i) {
if (vColInd[static_cast<std::size_t>(i)] == iGlobalCol) {
vLocalVal[static_cast<std::size_t>(i)] += cT.m_vVal;
break;
}
}
}
}
};

// Type aliases for common scalar types
using cCSRCommF = cCSRComm<float>;
using cCSRCommD = cCSRComm<double>;
using cCSRCommCF = cCSRComm<std::complex<float>>;
using cCSRCommCD = cCSRComm<std::complex<double>>;

}
Loading