scikit-learn/scikits/learn/datasets/_svmlight_format.cpp

339 lines
9.0 KiB
C++
Raw Normal View History

/*
2011-06-11 12:03:35 +08:00
* Authors: Mathieu Blondel <mathieu@mblondel.org>
* Lars Buitinck <L.J.Buitinck@uva.nl>
*
* License: Simple BSD
*
* This module implements _load_svmlight_format, a fast and memory efficient
* function to load the file format originally created for svmlight and now used
* by many other libraries, including libsvm.
*
* The function loads the file directly in a CSR sparse matrix without memory
* copying. The approach taken is to use 4 C++ vectors (data, indices, indptr
2011-06-11 01:04:10 +08:00
* and labels) and to incrementally feed them with elements. Ndarrays are then
2011-06-11 05:10:22 +08:00
* instantiated by PyArray_SimpleNewFromData, i.e., no memory is
2011-06-11 01:04:10 +08:00
* copied.
*
2011-06-11 12:54:12 +08:00
* Since the memory is not allocated by the ndarray, the ndarray doesn't own the
* memory and thus cannot deallocate it. To automatically deallocate memory, the
* technique described at http://blog.enthought.com/?p=62 is used. The main idea
* is to use an additional object that the ndarray does own and that will be
* responsible for deallocating the memory.
*/
#include <Python.h>
#include <numpy/arrayobject.h>
#include <fstream>
#include <sstream>
#include <stdexcept>
#include <string>
#include <vector>
/*
* A Python object responsible for memory management of our vectors.
*/
template <typename T>
struct VectorOwner {
2011-06-13 00:28:21 +08:00
// Inherit from the base Python object.
PyObject_HEAD
2011-06-13 00:28:21 +08:00
// The vector that VectorOwner is responsible for deallocating.
std::vector<T> v;
};
/*
2011-06-13 00:28:21 +08:00
* Deallocator template.
*/
template <typename T>
static void destroy_vector_owner(PyObject *self)
{
// Note: explicit call to destructor because of placement new in
// to_1d_array. memory management for VectorOwner is performed by Python.
2011-06-13 00:28:21 +08:00
// Compiler-generated destructor will release memory from vector member.
VectorOwner<T> &obj = *reinterpret_cast<VectorOwner<T> *>(self);
2011-06-12 23:54:29 +08:00
obj.~VectorOwner<T>();
self->ob_type->tp_free(self);
}
2011-06-13 00:28:21 +08:00
/*
* Since a template function can't have C linkage,
* we instantiate the template for the types "int" and "double"
* in the following two functions. These are used for the tp_dealloc
* attribute of the vector owner types further below.
*/
2011-06-11 19:20:16 +08:00
extern "C" {
static void destroy_int_vector(PyObject *self)
{
destroy_vector_owner<int>(self);
}
static void destroy_double_vector(PyObject *self)
{
destroy_vector_owner<double>(self);
}
2011-06-11 19:20:16 +08:00
}
/*
* Type objects for above.
*/
static PyTypeObject IntVOwnerType = { PyObject_HEAD_INIT(NULL) },
DoubleVOwnerType = { PyObject_HEAD_INIT(NULL) };
/*
* Set the fields of the owner type objects.
*/
static void init_type_objs()
{
IntVOwnerType.tp_flags = DoubleVOwnerType.tp_flags = Py_TPFLAGS_DEFAULT;
IntVOwnerType.tp_name = DoubleVOwnerType.tp_name = "deallocator";
IntVOwnerType.tp_doc = DoubleVOwnerType.tp_doc = "deallocator object";
IntVOwnerType.tp_new = DoubleVOwnerType.tp_new = PyType_GenericNew;
IntVOwnerType.tp_basicsize = sizeof(VectorOwner<int>);
DoubleVOwnerType.tp_basicsize = sizeof(VectorOwner<double>);
IntVOwnerType.tp_dealloc = destroy_int_vector;
DoubleVOwnerType.tp_dealloc = destroy_double_vector;
}
PyTypeObject &vector_owner_type(int typenum)
{
switch (typenum) {
case NPY_INT: return IntVOwnerType;
case NPY_DOUBLE: return DoubleVOwnerType;
}
throw std::logic_error("invalid argument to vector_owner_type");
}
/*
* Convert a C++ vector to a 1d-ndarray WITHOUT memory copying.
* Steals v's contents, leaving it empty.
* Throws an exception if an error occurs.
*/
template <typename T>
static PyObject *to_1d_array(std::vector<T> &v, int typenum)
{
npy_intp dims[1] = {v.size()};
// A C++ vector's elements are guaranteed to be in a contiguous array.
PyObject *arr = PyArray_SimpleNewFromData(1, dims, typenum, &v[0]);
try {
if (!arr)
throw std::bad_alloc();
VectorOwner<T> *owner = PyObject_New(VectorOwner<T>,
&vector_owner_type(typenum));
if (!owner)
throw std::bad_alloc();
2011-06-13 00:28:21 +08:00
// Transfer ownership of v's contents to the VectorOwner.
// Note: placement new.
new (&owner->v) std::vector<T>();
owner->v.swap(v);
PyArray_BASE(arr) = (PyObject *)owner;
return arr;
} catch (std::exception const &e) {
2011-06-13 00:28:21 +08:00
// Let's assume the Python exception is already set correctly.
Py_XDECREF(arr);
throw;
}
}
static PyObject *to_csr(std::vector<double> &data,
std::vector<int> &indices,
std::vector<int> &indptr,
std::vector<double> &labels)
{
// We could do with a smart pointer to Python objects here.
std::exception const *exc = 0;
PyObject *data_arr = 0,
*indices_arr = 0,
*indptr_arr = 0,
*labels_arr = 0,
*ret_tuple = 0;
try {
data_arr = to_1d_array(data, NPY_DOUBLE);
indices_arr = to_1d_array(indices, NPY_INT);
indptr_arr = to_1d_array(indptr, NPY_INT);
labels_arr = to_1d_array(labels, NPY_DOUBLE);
ret_tuple = Py_BuildValue("OOOO",
data_arr, indices_arr,
indptr_arr, labels_arr);
} catch (std::exception const &e) {
exc = &e;
}
// Py_BuildValue increases the reference count of each array,
// so we need to decrease it before returning the tuple,
// regardless of error status.
Py_XDECREF(data_arr);
Py_XDECREF(indices_arr);
Py_XDECREF(indptr_arr);
Py_XDECREF(labels_arr);
if (exc)
throw *exc;
return ret_tuple;
}
/*
* Parsing.
*/
class SyntaxError : public std::runtime_error {
public:
SyntaxError(std::string const &msg)
: std::runtime_error(msg + " in SVMlight/libSVM file")
{
}
};
/*
* Parse single line. Throws exception on failure.
*/
void parse_line(const std::string& line,
std::vector<double> &data,
std::vector<int> &indices,
std::vector<int> &indptr,
std::vector<double> &labels)
{
2011-06-11 05:10:22 +08:00
if (line.length() == 0)
throw SyntaxError("empty line");
if (line[0] == '#')
return;
2011-06-21 17:31:46 +08:00
// FIXME: we shouldn't be parsing line-by-line.
// Also, we might catch more syntax errors with failbit.
std::istringstream in(line);
in.exceptions(std::ios::badbit);
2011-06-21 17:31:46 +08:00
double y;
if (!(in >> y))
throw SyntaxError("non-numeric or missing label");
2011-06-11 05:10:22 +08:00
labels.push_back(y);
indptr.push_back(data.size());
char c;
double x;
unsigned idx;
2011-06-21 17:31:46 +08:00
while (in >> idx >> c >> x) {
if (c != ':')
throw SyntaxError(std::string("expected ':', got '") + c + "'");
indices.push_back(int(idx));
data.push_back(x);
}
}
/*
* Parse entire file. Throws exception on failure.
*/
static void parse_file(char const *file_path,
size_t buffer_size,
std::vector<double> &data,
std::vector<int> &indices,
std::vector<int> &indptr,
std::vector<double> &labels)
{
2011-06-11 05:10:22 +08:00
std::vector<char> buffer(buffer_size);
std::ifstream file_stream;
file_stream.exceptions(std::ios::badbit);
2011-06-12 23:54:29 +08:00
file_stream.rdbuf()->pubsetbuf(&buffer[0], buffer_size);
file_stream.open(file_path);
2011-06-21 16:57:21 +08:00
if (!file_stream)
throw std::ios_base::failure("File doesn't exist!");
std::string line;
while (std::getline(file_stream, line))
parse_line(line, data, indices, indptr, labels);
indptr.push_back(data.size());
}
static const char load_svmlight_file_doc[] =
2011-06-11 05:10:22 +08:00
"Load file in svmlight format and return a CSR.";
extern "C" {
static PyObject *load_svmlight_file(PyObject *self, PyObject *args)
{
try {
2011-06-13 00:28:21 +08:00
// Read function arguments.
char const *file_path;
int buffer_mb;
if (!PyArg_ParseTuple(args, "si", &file_path, &buffer_mb))
return 0;
buffer_mb = std::max(buffer_mb, 1);
size_t buffer_size = buffer_mb * 1024 * 1024;
std::vector<double> data, labels;
std::vector<int> indices, indptr;
parse_file(file_path, buffer_size, data, indices, indptr, labels);
return to_csr(data, indices, indptr, labels);
} catch (SyntaxError const &e) {
PyErr_SetString(PyExc_ValueError, e.what());
return 0;
} catch (std::bad_alloc const &e) {
PyErr_SetString(PyExc_MemoryError, e.what());
return 0;
} catch (std::ios_base::failure const &e) {
PyErr_SetString(PyExc_IOError, e.what());
return 0;
} catch (std::exception const &e) {
std::string msg("error in SVMlight/libSVM reader: ");
msg += e.what();
PyErr_SetString(PyExc_RuntimeError, msg.c_str());
return 0;
}
}
}
/*
* Python module setup.
*/
static PyMethodDef svmlight_format_methods[] = {
{"_load_svmlight_file", load_svmlight_file,
METH_VARARGS, load_svmlight_file_doc},
{NULL, NULL, 0, NULL}
};
2011-06-11 19:20:16 +08:00
static const char svmlight_format_doc[] =
"Loader for svmlight / libsvm datasets - C++ helper routines";
extern "C" {
PyMODINIT_FUNC init_svmlight_format(void)
{
_import_array();
init_type_objs();
if (PyType_Ready(&DoubleVOwnerType) < 0
|| PyType_Ready(&IntVOwnerType) < 0)
return;
Py_InitModule3("_svmlight_format",
svmlight_format_methods,
svmlight_format_doc);
}
}