Add printoptions (#3333)

This commit is contained in:
Christophe Prat
2026-04-01 22:24:48 -07:00
committed by GitHub
parent 80a1c206f9
commit befe42d303
9 changed files with 174 additions and 13 deletions
+2 -1
View File
@@ -27,7 +27,8 @@ nanobind_add_module(
${CMAKE_CURRENT_SOURCE_DIR}/linalg.cpp
${CMAKE_CURRENT_SOURCE_DIR}/constants.cpp
${CMAKE_CURRENT_SOURCE_DIR}/trees.cpp
${CMAKE_CURRENT_SOURCE_DIR}/utils.cpp)
${CMAKE_CURRENT_SOURCE_DIR}/utils.cpp
${CMAKE_CURRENT_SOURCE_DIR}/print.cpp)
if(MLX_BUILD_PYTHON_STUBS)
nanobind_add_stub(
+1 -3
View File
@@ -12,6 +12,7 @@
#include <nanobind/typing.h>
#include "mlx/backend/metal/metal.h"
#include "mlx/utils.h"
#include "python/src/buffer.h"
#include "python/src/convert.h"
#include "python/src/indexing.h"
@@ -97,9 +98,6 @@ class ArrayPythonIterator {
};
void init_array(nb::module_& m) {
// Set Python print formatting options
mx::get_global_formatter().capitalize_bool = true;
// Types
nb::class_<mx::Dtype>(
m,
+2
View File
@@ -23,6 +23,7 @@ void init_constants(nb::module_&);
void init_fast(nb::module_&);
void init_distributed(nb::module_&);
void init_export(nb::module_&);
void init_print(nb::module_&);
NB_MODULE(core, m) {
m.doc() = "mlx: A framework for machine learning on Apple silicon.";
@@ -46,6 +47,7 @@ NB_MODULE(core, m) {
init_fast(m);
init_distributed(m);
init_export(m);
init_print(m);
m.attr("__version__") = mx::version();
}
+86
View File
@@ -0,0 +1,86 @@
#include <cstdint>
#include <cstring>
#include <sstream>
#include <nanobind/typing.h>
#include "mlx/utils.h"
#include "python/src/utils.h"
#include "mlx/mlx.h"
namespace mx = mlx::core;
namespace nb = nanobind;
using namespace nb::literals;
struct PrintOptionsContext {
mx::PrintOptions old_options;
mx::PrintOptions new_options;
PrintOptionsContext(mx::PrintOptions p) : new_options(p) {}
PrintOptionsContext& enter() {
old_options = mx::get_global_formatter().format_options;
mx::set_printoptions(new_options);
return *this;
}
void exit(nb::args) {
mx::set_printoptions(old_options);
}
};
void init_print(nb::module_& m) {
// Set Python print formatting options
mx::get_global_formatter().capitalize_bool = true;
// Expose printing options to Python: allow setting global precision.
nb::class_<mx::PrintOptions>(m, "PrintOptions")
.def(nb::init<int>(), "precision"_a = -1)
.def_rw("precision", &mx::PrintOptions::precision);
m.def(
"set_printoptions",
[](int precision) { mx::set_printoptions({precision}); },
"precision"_a = mx::get_global_formatter().format_options.precision,
R"pbdoc(
Set global printing precision for array formatting.
Example:
>>> print(x) # Uses default precision
>>> mx.set_printoptions(precision=3):
>>> print(x) # Uses precision of 3
>>> print(x) # Uses precision of 3 (again)
Args:
precision (int): Number of decimal places.
)pbdoc");
m.def(
"get_printoptions",
[]() { return mx::get_global_formatter().format_options; },
R"pbdoc(
Get global printing precision for array formatting.
Returns:
PrintOptions: The format options used for printing arrays.
)pbdoc");
nb::class_<PrintOptionsContext>(m, "_PrintOptionsContext")
.def(nb::init<mx::PrintOptions>())
.def("__enter__", &PrintOptionsContext::enter)
.def("__exit__", &PrintOptionsContext::exit);
m.def(
"printoptions",
[](int precision) { return PrintOptionsContext({precision}); },
"precision"_a = mx::get_global_formatter().format_options.precision,
R"pbdoc(
Context manager for setting print options temporarily.
Example:
>>> print(x) # Uses default precision
>>> with mx.printoptions(precision=3):
>>> print(x) # Uses precision of 3
>>> print(x) # Back to default precision
Args:
precision (int): Number of decimal places. Use -1 for default
)pbdoc");
}