DDC 0.16.0
Loading...
Searching...
No Matches
save_npy.cpp
1// Copyright (C) The DDC development team, see COPYRIGHT.md file
2//
3// SPDX-License-Identifier: MIT
4
5#include <array>
6#include <bit>
7#include <cstddef>
8#include <cstdint>
9#include <filesystem>
10#include <fstream>
11#include <functional>
12#include <numeric>
13#include <stdexcept>
14#include <string>
15#include <utility>
16#include <vector>
17
18#include "save_npy.hpp"
19
20namespace ddc::detail {
21
22NpyByteOrder get_byte_order(std::size_t const itemsize) noexcept
23{
24 if (itemsize == 1) {
25 return NpyByteOrder::not_applicable;
26 }
27
28 if (std::endian::native == std::endian::little) {
29 return NpyByteOrder::little_endian;
30 }
31
32 if (std::endian::native == std::endian::big) {
33 return NpyByteOrder::big_endian;
34 }
35
36 return NpyByteOrder::not_applicable;
37}
38
39void write_le(std::ostream& os, std::uint16_t const value_u16)
40{
41 constexpr unsigned int mask = 0xFFU;
42
43 unsigned int const value_u = value_u16;
44
45 std::array<char, 2> bytes;
46 bytes[0] = value_u & mask;
47 bytes[1] = (value_u >> 8U) & mask;
48
49 os.write(bytes.data(), sizeof(value_u16));
50}
51
52std::string NpyDtype::str() const
53{
54 return std::string(1, static_cast<char>(byte_order)) + static_cast<char>(kind)
55 + std::to_string(itemsize);
56}
57
58// See specification at https://numpy.org/neps/nep-0001-npy-format.html#format-specification-version-1-0
59void save_npy(std::ostream& os, NpyArrayView const& view)
60{
61 // Build shape string: (d0, d1, ..., dN,)
62 std::string shape_str = "(";
63 for (std::size_t const ext : view.shape) {
64 shape_str += std::to_string(ext);
65 shape_str += ", ";
66 }
67 shape_str += ")";
68
69 std::string const header_dict
70 = std::string("{'descr': '") + view.dtype.str() + "', 'fortran_order': "
71 + (view.fortran_order ? "True" : "False") + ", 'shape': " + shape_str + ", }";
72
73 // Pad header to a multiple of 16
74 std::size_t const non_padded_header_len = header_dict.size() + 1;
75 // magic(6) + major(1) + minor(1) + header_len(2) + header
76 std::size_t const alignment = 16;
77 std::size_t const remainder = (6 + 1 + 1 + 2 + non_padded_header_len) % alignment;
78 std::size_t const padding = (alignment - remainder) % alignment;
79 if (!std::in_range<std::uint16_t>(non_padded_header_len + padding)) {
80 throw std::runtime_error("save_npy: header too large for npy v1.0.");
81 }
82 auto const header_len = static_cast<std::uint16_t>(non_padded_header_len + padding);
83
84 // magic string
85 os.write("\x93NUMPY", 6);
86 // major version
87 os.put(1);
88 // minor version
89 os.put(0);
90 // header length in little-endian
91 write_le(os, header_len);
92 // header + padding + newline
93 os.write(header_dict.data(), header_dict.size());
94 os.write(" ", padding);
95 os.put('\n');
96
97 // Raw data
98 std::size_t const n_elems
99 = std::accumulate(view.shape.begin(), view.shape.end(), 1ULL, std::multiplies<> {});
100 os.write(reinterpret_cast<char const*>(view.data), n_elems * view.dtype.itemsize);
101}
102
103void save_npy(std::filesystem::path const& filename, NpyArrayView const& view)
104{
105 std::ofstream file(filename, std::ios::binary);
106 file.exceptions(std::ios::failbit | std::ios::badbit);
107
108 save_npy(file, view);
109}
110
111} // namespace ddc::detail
The top-level namespace of DDC.