// ============================================================
// dlarray.cpp --- dlarray 精简 C++ 实现 + 使用示例
// 编译:g++ -std=c++17 -O2 dlarray.cpp -o dlarray
// 运行:./dlarray
// ============================================================
#include <vector>
#include <string>
#include <memory>
#include <functional>
#include <stdexcept>
#include <cmath>
#include <algorithm>
#include <numeric>
#include <unordered_map>
#include <iostream>
#include <cstddef>
// ============================================================
// 格式标签工具
// ============================================================
namespace dlfmt {
inline const std::string kOrder = "SCBTU";
inline bool isValidLabel(char c) {
return kOrder.find(c) != std::string::npos;
}
// 按 SCBTU 顺序重排(S 可重复,C/B/T 最多一次)
inline std::string normalize(const std::string& fmt) {
std::string out;
for (char c : fmt)
if (c == 'S') out += 'S';
for (char target : std::string("CBTU")) {
for (char c : fmt)
if (c == target) { out += target; break; }
}
return out;
}
inline void validate(const std::string& fmt) {
std::unordered_map<char, int> cnt;
for (char c : fmt) {
if (!isValidLabel(c))
throw std::invalid_argument(
"Invalid dlarray format label: " + std::string(1, c));
cnt[c]++;
}
for (char c : std::string("CBT")) {
if (cnt[c] > 1)
throw std::invalid_argument(
"Label '" + std::string(1, c) +
"' can appear at most once in dlarray format");
}
}
} // namespace dlfmt
// ============================================================
// DLArray 主类
// ============================================================
class DLArray {
public:
using DataType = float;
using Shape = std::vector<size_t>;
using DataVec = std::vector<DataType>;
// ---------------- 构造 ----------------
DLArray() : m_requiresGrad(false) {}
explicit DLArray(const Shape& shape)
: m_shape(shape), m_data(shapeProduct(shape), 0.0f),
m_requiresGrad(false) {}
DLArray(const Shape& shape, const std::string& fmt)
: m_shape(shape), m_format(dlfmt::normalize(fmt)),
m_data(shapeProduct(shape), 0.0f), m_requiresGrad(false) {
dlfmt::validate(m_format);
if (m_format.size() > m_shape.size())
throw std::invalid_argument("Format label count exceeds dimension count");
}
DLArray(const DataVec& data, const Shape& shape, const std::string& fmt = "")
: m_shape(shape), m_format(dlfmt::normalize(fmt)), m_data(data),
m_requiresGrad(false) {
if (m_data.size() != shapeProduct(shape))
throw std::invalid_argument("Data size does not match shape");
dlfmt::validate(m_format);
}
static DLArray scalar(DataType v) {
DLArray a({1}); a.m_data[0] = v; return a;
}
// ---------------- 数据访问 ----------------
const DataType* data() const { return m_data.data(); }
DataType* data() { return m_data.data(); }
const DataVec& vec() const { return m_data; }
DataVec& vec() { return m_data; }
const Shape& shape() const { return m_shape; }
const std::string& format() const { return m_format; }
size_t numel() const { return m_data.size(); }
size_t ndims() const { return m_shape.size(); }
bool isFormatted() const { return !m_format.empty(); }
DataVec extractdata() const { return m_data; }
std::string dims() const { return m_format; }
// stripdims:移除格式标签
DLArray stripdims() const {
DLArray r(*this); r.m_format.clear(); return r;
}
// finddim:返回 1 起始的维度索引
std::vector<size_t> finddim(const std::string& labels) const {
std::vector<size_t> idx;
for (char lab : labels) {
bool found = false;
for (size_t i = 0; i < m_format.size(); ++i)
if (m_format[i] == lab) { idx.push_back(i + 1); found = true; break; }
if (!found)
throw std::runtime_error(
"Label '" + std::string(1, lab) + "' not found in format");
}
return idx;
}
// ---------------- 自动微分接口 ----------------
void setRequiresGrad(bool v) { m_requiresGrad = v; }
bool requiresGrad() const { return m_requiresGrad; }
DLArray& grad() {
if (!m_grad) m_grad = std::make_shared<DLArray>(m_shape);
return *m_grad;
}
const DLArray& grad() const {
if (!m_grad) m_grad = std::make_shared<DLArray>(m_shape);
return *m_grad;
}
void zeroGrad() {
if (m_grad && m_grad->m_data.size() == m_data.size())
std::fill(m_grad->m_data.begin(), m_grad->m_data.end(), 0.0f);
}
void backward(const DLArray& gradOut) {
if (m_backward) m_backward(gradOut);
}
// ---------------- 算术运算(含自动微分) ----------------
DLArray operator+(DLArray& other) { return binaryOp(other, "add"); }
DLArray operator-(DLArray& other) { return binaryOp(other, "sub"); }
DLArray operator*(DLArray& other) { return binaryOp(other, "mul"); }
DLArray operator/(DLArray& other) { return binaryOp(other, "div"); }
DLArray operator-() {
DLArray r(*this);
for (auto& v : r.m_data) v = -v;
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this](const DLArray& g) {
DLArray ng(g);
for (auto& v : ng.m_data) v = -v;
m_accumulateGrad(ng);
if (m_backward) m_backward(ng);
};
}
return r;
}
DLArray pow(DataType e) {
DLArray r(*this);
for (auto& v : r.m_data) v = std::pow(v, e);
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this, e](const DLArray& g) {
DLArray grad(m_shape);
for (size_t i = 0; i < m_data.size(); ++i)
grad.m_data[i] = g.m_data[i] * e *
std::pow(m_data[i], e - 1.0f);
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return r;
}
// ---------------- 矩阵乘法 ----------------
DLArray mtimes(DLArray& other) {
if (m_shape.size() != 2 || other.m_shape.size() != 2)
throw std::runtime_error("mtimes: only 2-D supported");
if (m_shape[1] != other.m_shape[0])
throw std::runtime_error("mtimes: inner dimension mismatch");
size_t M = m_shape[0], K = m_shape[1], N = other.m_shape[1];
DLArray result({M, N}, m_format);
for (size_t i = 0; i < M; ++i)
for (size_t j = 0; j < N; ++j) {
DataType s = 0;
for (size_t k = 0; k < K; ++k)
s += m_data[i * K + k] * other.m_data[k * N + j];
result.m_data[i * N + j] = s;
}
if (m_requiresGrad || other.m_requiresGrad) {
result.m_requiresGrad = true;
result.m_backward = [this, &other, M, K, N](const DLArray& g) {
if (m_requiresGrad) {
DLArray dA({M, K});
for (size_t i = 0; i < M; ++i)
for (size_t k = 0; k < K; ++k) {
DataType s = 0;
for (size_t j = 0; j < N; ++j)
s += g.m_data[i * N + j] * other.m_data[k * N + j];
dA.m_data[i * K + k] = s;
}
m_accumulateGrad(dA);
if (m_backward) m_backward(dA);
}
if (other.m_requiresGrad) {
DLArray dB({K, N});
for (size_t k = 0; k < K; ++k)
for (size_t j = 0; j < N; ++j) {
DataType s = 0;
for (size_t i = 0; i < M; ++i)
s += m_data[i * K + k] * g.m_data[i * N + j];
dB.m_data[k * N + j] = s;
}
other.m_accumulateGrad(dB);
if (other.m_backward) other.m_backward(dB);
}
};
}
return result;
}
// ---------------- 深度学习运算 ----------------
DLArray softmax(int dim = -1) {
DLArray result(*this);
if (dim < 0) dim = static_cast<int>(m_shape.size()) - 1;
size_t d = static_cast<size_t>(dim);
size_t stride = 1;
for (size_t i = d + 1; i < m_shape.size(); ++i) stride *= m_shape[i];
size_t outer = m_data.size() / (m_shape[d] * stride);
for (size_t o = 0; o < outer; ++o)
for (size_t s = 0; s < stride; ++s) {
size_t base = o * m_shape[d] * stride + s;
DataType mx = -1e30f;
for (size_t k = 0; k < m_shape[d]; ++k)
mx = std::max(mx, m_data[base + k * stride]);
DataType sm = 0;
for (size_t k = 0; k < m_shape[d]; ++k) {
DataType e = std::exp(m_data[base + k * stride] - mx);
result.m_data[base + k * stride] = e;
sm += e;
}
for (size_t k = 0; k < m_shape[d]; ++k)
result.m_data[base + k * stride] /= sm;
}
if (m_requiresGrad) {
result.m_requiresGrad = true;
result.m_backward = [this, result, d, stride, outer](const DLArray& g) {
DLArray grad(m_shape);
for (size_t o = 0; o < outer; ++o)
for (size_t s = 0; s < stride; ++s) {
size_t base = o * m_shape[d] * stride + s;
DataType dot = 0;
for (size_t k = 0; k < m_shape[d]; ++k)
dot += g.m_data[base + k * stride] *
result.m_data[base + k * stride];
for (size_t k = 0; k < m_shape[d]; ++k) {
size_t pos = base + k * stride;
grad.m_data[pos] = result.m_data[pos] *
(g.m_data[pos] - dot);
}
}
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return result;
}
DLArray sigmoid() {
DLArray r(*this);
for (auto& v : r.m_data) v = 1.0f / (1.0f + std::exp(-v));
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this, r](const DLArray& g) {
DLArray grad(m_shape);
for (size_t i = 0; i < m_data.size(); ++i)
grad.m_data[i] = g.m_data[i] * r.m_data[i] * (1.0f - r.m_data[i]);
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return r;
}
DLArray relu() {
DLArray r(*this);
for (auto& v : r.m_data) v = std::max(0.0f, v);
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this](const DLArray& g) {
DLArray grad(m_shape);
for (size_t i = 0; i < m_data.size(); ++i)
grad.m_data[i] = m_data[i] > 0.0f ? g.m_data[i] : 0.0f;
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return r;
}
DLArray sum(int dim = -1) {
if (dim < 0) {
DLArray r({1});
r.m_data[0] = std::accumulate(m_data.begin(), m_data.end(), 0.0f);
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this](const DLArray& g) {
DLArray grad(m_shape);
std::fill(grad.m_data.begin(), grad.m_data.end(), g.m_data[0]);
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return r;
}
size_t d = static_cast<size_t>(dim);
size_t stride = 1;
for (size_t i = d + 1; i < m_shape.size(); ++i) stride *= m_shape[i];
size_t outer = m_data.size() / (m_shape[d] * stride);
Shape outShape = m_shape;
outShape[d] = 1;
DLArray r(outShape, m_format);
for (size_t o = 0; o < outer; ++o)
for (size_t s = 0; s < stride; ++s) {
size_t base = o * m_shape[d] * stride + s;
DataType acc = 0;
for (size_t k = 0; k < m_shape[d]; ++k)
acc += m_data[base + k * stride];
r.m_data[o * stride + s] = acc;
}
if (m_requiresGrad) {
r.m_requiresGrad = true;
r.m_backward = [this, d, stride, outer](const DLArray& g) {
DLArray grad(m_shape);
for (size_t o = 0; o < outer; ++o)
for (size_t s = 0; s < stride; ++s) {
DataType gv = g.m_data[o * stride + s];
size_t base = o * m_shape[d] * stride + s;
for (size_t k = 0; k < m_shape[d]; ++k)
grad.m_data[base + k * stride] = gv;
}
m_accumulateGrad(grad);
if (m_backward) m_backward(grad);
};
}
return r;
}
DLArray mean(int dim = -1) {
DLArray s = sum(dim);
if (dim < 0) s.m_data[0] /= static_cast<DataType>(m_data.size());
else s.m_data[0] /= static_cast<DataType>(m_shape[(size_t)dim]);
return s;
}
private:
Shape m_shape;
std::string m_format;
DataVec m_data;
bool m_requiresGrad;
mutable std::shared_ptr<DLArray> m_grad;
std::function<void(const DLArray&)> m_backward;
// ---------------- 辅助函数 ----------------
static size_t shapeProduct(const Shape& s) {
size_t p = 1;
for (auto d : s) p *= d;
return p;
}
void m_accumulateGrad(const DLArray& g) const {
if (!m_grad) m_grad = std::make_shared<DLArray>(m_shape);
if (m_grad->m_data.size() != g.m_data.size())
throw std::runtime_error("Gradient size mismatch");
for (size_t i = 0; i < m_grad->m_data.size(); ++i)
m_grad->m_data[i] += g.m_data[i];
}
static std::string combineFormat(const std::string& fa, const std::string& fb) {
size_t sCount = std::max(
(size_t)std::count(fa.begin(), fa.end(), 'S'),
(size_t)std::count(fb.begin(), fb.end(), 'S'));
std::string out(sCount, 'S');
for (char c : std::string("CBTU"))
if (fa.find(c) != std::string::npos || fb.find(c) != std::string::npos)
out += c;
return out;
}
static size_t labelDim(const Shape& shape, const std::string& fmt, char label) {
for (size_t i = 0; i < fmt.size() && i < shape.size(); ++i)
if (fmt[i] == label) return shape[i];
return 1;
}
static Shape broadcastShape(const Shape& a, const std::string& fa,
const Shape& b, const std::string& fb,
std::string& outFmt) {
if (fa.empty() || fb.empty()) {
outFmt = fa.empty() ? fb : fa;
size_t n = std::max(a.size(), b.size());
Shape out(n);
for (size_t i = 0; i < n; ++i) {
size_t da = (i < a.size()) ? a[a.size() - 1 - i] : 1;
size_t db = (i < b.size()) ? b[b.size() - 1 - i] : 1;
if (da != db && da != 1 && db != 1)
throw std::runtime_error("Broadcast failed");
out[n - 1 - i] = std::max(da, db);
}
return out;
}
outFmt = combineFormat(fa, fb);
Shape out(outFmt.size());
for (size_t i = 0; i < outFmt.size(); ++i) {
char lab = outFmt[i];
size_t da = labelDim(a, fa, lab);
size_t db = labelDim(b, fb, lab);
if (da != db && da != 1 && db != 1)
throw std::runtime_error("Broadcast label mismatch");
out[i] = std::max(da, db);
}
return out;
}
static DataVec broadcastData(const DataVec& src, const Shape& srcShape,
const std::string& srcFmt,
const Shape& dstShape, const std::string& dstFmt) {
DataVec dst(shapeProduct(dstShape));
for (size_t i = 0; i < dst.size(); ++i) {
std::vector<size_t> dstIdx(dstShape.size());
size_t tmp = i;
for (int d = (int)dstShape.size() - 1; d >= 0; --d) {
dstIdx[d] = tmp % dstShape[d];
tmp /= dstShape[d];
}
std::vector<size_t> srcIdx(srcShape.size(), 0);
if (srcFmt.empty() || dstFmt.empty()) {
size_t offset = dstShape.size() - srcShape.size();
for (size_t d = 0; d < srcShape.size(); ++d)
srcIdx[d] = (srcShape[d] == 1) ? 0 : dstIdx[d + offset];
} else {
for (size_t sd = 0; sd < srcFmt.size(); ++sd) {
char lab = srcFmt[sd];
for (size_t dd = 0; dd < dstFmt.size(); ++dd)
if (dstFmt[dd] == lab) {
srcIdx[sd] = (srcShape[sd] == 1) ? 0 : dstIdx[dd];
break;
}
}
}
size_t lin = 0, stride = 1;
for (int d = (int)srcShape.size() - 1; d >= 0; --d) {
lin += srcIdx[d] * stride;
stride *= srcShape[d];
}
dst[i] = src[lin];
}
return dst;
}
DLArray binaryOp(DLArray& other, const std::string& opName) {
std::string outFmt;
Shape outShape = broadcastShape(m_shape, m_format,
other.m_shape, other.m_format, outFmt);
DataVec a = broadcastData(m_data, m_shape, m_format, outShape, outFmt);
DataVec b = broadcastData(other.m_data, other.m_shape, other.m_format,
outShape, outFmt);
DLArray result(outShape, outFmt);
size_t total = result.m_data.size();
for (size_t i = 0; i < total; ++i) {
if (opName == "add") result.m_data[i] = a[i] + b[i];
else if (opName == "sub") result.m_data[i] = a[i] - b[i];
else if (opName == "mul") result.m_data[i] = a[i] * b[i];
else if (opName == "div") {
if (b[i] == 0.0f) throw std::runtime_error("Division by zero");
result.m_data[i] = a[i] / b[i];
}
}
if (m_requiresGrad || other.m_requiresGrad) {
result.m_requiresGrad = true;
result.m_backward = [this, &other, a, b, opName](const DLArray& g) {
if (m_requiresGrad) {
DLArray ga(m_shape);
for (size_t i = 0; i < ga.m_data.size(); ++i) {
size_t j = i % g.m_data.size();
if (opName == "add") ga.m_data[i] = g.m_data[j];
else if (opName == "sub") ga.m_data[i] = g.m_data[j];
else if (opName == "mul") ga.m_data[i] = g.m_data[j] * b[j];
else if (opName == "div") ga.m_data[i] = g.m_data[j] / b[j];
}
m_accumulateGrad(ga);
if (m_backward) m_backward(ga);
}
if (other.m_requiresGrad) {
DLArray gb(other.m_shape);
for (size_t i = 0; i < gb.m_data.size(); ++i) {
size_t j = i % g.m_data.size();
if (opName == "add") gb.m_data[i] = g.m_data[j];
else if (opName == "sub") gb.m_data[i] = -g.m_data[j];
else if (opName == "mul") gb.m_data[i] = g.m_data[j] * a[j];
else if (opName == "div")
gb.m_data[i] = -g.m_data[j] * a[j] / (b[j] * b[j]);
}
other.m_accumulateGrad(gb);
if (other.m_backward) other.m_backward(gb);
}
};
}
return result;
}
};
// ============================================================
// main 示例
// ============================================================
int main() {
// ---- 1. 创建格式化 dlarray ----
DLArray X({2, 3, 2}, "SCB");
DLArray Y({3}, "C");
for (size_t i = 0; i < X.numel(); ++i)
X.data()[i] = static_cast<float>(i + 1);
for (size_t i = 0; i < Y.numel(); ++i)
Y.data()[i] = static_cast<float>(i + 1);
std::cout << "X format: " << X.format()
<< ", numel: " << X.numel() << "\n";
std::cout << "Y format: " << Y.format() << "\n";
// ---- 2. 格式标签操作 ----
auto sDims = X.finddim("S");
std::cout << "S dimension index: " << sDims[0] << "\n";
DLArray Xs = X.stripdims();
std::cout << "After stripdims, formatted? " << Xs.isFormatted() << "\n";
// ---- 3. 算术运算(按标签广播) ----
DLArray Z = X * Y;
std::cout << "Z format: " << Z.format() << "\n";
// ---- 4. 自动微分 ----
DLArray A({2, 3}, "SC");
A.setRequiresGrad(true);
for (size_t i = 0; i < A.numel(); ++i)
A.data()[i] = static_cast<float>(i + 1);
DLArray B = A * A; // B = A .^ 2
DLArray C = B.sum(); // C = sum(A .^ 2)
DLArray gradOut({1});
gradOut.data()[0] = 1.0f;
C.backward(gradOut);
std::cout << "d(sum(A^2))/dA = 2*A:\n";
for (size_t i = 0; i < A.numel(); ++i)
std::cout << A.grad().data()[i] << " ";
std::cout << "\n";
// ---- 5. 深度学习运算 ----
DLArray logits({3, 1}, "CB");
logits.data()[0] = 1.0f;
logits.data()[1] = 2.0f;
logits.data()[2] = 3.0f;
DLArray probs = logits.softmax(0);
std::cout << "Softmax: ";
for (size_t i = 0; i < probs.numel(); ++i)
std::cout << probs.data()[i] << " ";
std::cout << "\n";
// ---- 6. extractdata ----
auto raw = probs.extractdata();
std::cout << "Extracted: ";
for (auto v : raw) std::cout << v << " ";
std::cout << "\n";
return 0;
}