How to remove offsets in 3D projection plot?

Viewed 126

I realized a slight "misalignment" of a plot I'm making in 3D with matplotlib. Here is an MWE:

import numpy as np
from matplotlib import pyplot as plt
figure = plt.figure(figsize=(8,10.7))
ax = plt.gca(projection='3d')
ax.plot_surface(np.array([[0, 0], [30, 30]]),
                np.array([[10, 10], [10, 10]]),
                np.array([[10, 20], [10, 20]]), 
                rstride=1, cstride=1
)
ax.plot_surface(np.array([[0, 0], [30, 30]]),
                np.array([[20, 20], [20, 20]]),
                np.array([[10, 20], [10, 20]]), 
                rstride=1, cstride=1
)
plt.show()
plt.close()

Clearly, the bins are not correctly centered, as the surfaces seem to start at 10.5 and end at 20.5 instead of 10 and 20 sharply. How could one achieve the latter?

enter image description here

EDIT: I'm afraid that there is an issue with the suggested answer. The x-axis does not have a solid black line, as is the case by default:

enter image description here

When I take out the suggested wrapping, I get:

enter image description here

Unfortunately, when I take out the stuff that I'm plotting, this issue is not reproducible in a Jupyter notebook, but nevertheless, I was wondering about whether you might be able to point out to me what I'd have to do so that in my case, the x-axis has a black line again?

1 Answers

This is caused by matplotlib's processing of 3D Axis coordinates. It deliberately shifts mins and maxs to create some artificial padding:

axis3d.py#L190-L194

class Axis(maxis.XAxis):
   ...
   def _get_coord_info(self, renderer):
       ...
       # Add a small offset between min/max point and the edge of the plot
       deltas = (maxs - mins) / 12
       mins -= 0.25 * deltas
       maxs += 0.25 * deltas
       ...
       return mins, maxs, centers, deltas, bounds_proj, highs

As of v3.5.1, there is no parameter to control this behavior.

However, we can use functools.wraps to create a wrapper around Axis._get_coord_info that unshifts mins and maxs. To prevent this wrapper from unshifting multiple times (e.g., when rerunning its Jupyter cell), track the wrapper state via an _unpadded attribute:

from functools import wraps

def unpad(f): # where f will be Axis._get_coord_info
    @wraps(f)
    def wrapper(*args, **kwargs):
        mins, maxs, centers, deltas, bounds_proj, highs = f(*args, **kwargs)
        mins += 0.25 * deltas # undo original subtraction
        maxs -= 0.25 * deltas # undo original addition
        return mins, maxs, centers, deltas, bounds_proj, highs

    if getattr(f, '_unpadded', False): # bypass if already unpadded
        return f
    else:
        wrapper._unpadded = True # mark as unpadded
        return wrapper

Apply the unpad wrapper before plotting:

import numpy as np
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d.axis3d import Axis

X = np.array([[0, 0], [30, 30]])
Z = np.array([[10, 20], [10, 20]])
y1, y2, y3 = 10, 16, 20

# wrap Axis._get_coord_info with our unpadded version
Axis._get_coord_info = unpad(Axis._get_coord_info)

fig, ax = plt.subplots(figsize=(8, 10.7), subplot_kw={'projection': '3d'})
ax.plot_surface(X, np.tile(y1, (2, 2)), Z, rstride=1, cstride=1)
ax.plot_surface(X, np.tile(y2, (2, 2)), Z, rstride=1, cstride=1)
ax.plot_surface(X, np.tile(y3, (2, 2)), Z, rstride=1, cstride=1)

plt.show()

unpadded 3d plot

Related