C++代码实现MATLAB中的dlarray函数功能

cpp 复制代码
// ============================================================
//  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;
}
相关推荐
一水鉴天4 小时前
映射、哈希表与哈斯图:计算机科学的三种基线 20261003(元宝)
开发语言·人工智能
Frank_refuel4 小时前
C++11之一场名为“搬家”的 C++ 之旅
开发语言·c++
Wang's Blog5 小时前
Java 项目实战: 外卖平台优化-Nginx配置文件结构与块层级
java·开发语言·nginx
2601_962071575 小时前
类变量和全局变量的查找路径有什么区别?
开发语言·python
hetao17338376 小时前
2026-10-03~04 hetao1733837 的刷题记录
c++·算法
\光辉岁月/7 小时前
5.java-数组
java·开发语言
Java后端的Ai之路7 小时前
Python进阶探索29_eval内置函数
开发语言·python·探索·eval·内置函数
谢亮_vipxieliang7 小时前
Spring 事务失效的常见场景
java·开发语言·数据库·spring boot
Zootopia6267 小时前
飞行力学知识梳理1|飞行性能与稳定性
人工智能·python·算法·机器学习·无人机·学习方法·信息与通信