eq_hist_new / eq_hist_new.py
christopher's picture
Create eq_hist_new.py
8758f90 verified
Raw
History Blame Contribute Delete
10.2 kB
# Would like to actually count the pixels of each color to see how that histogram turns out...
from bokeh.models import ColumnDataSource, Div, LinearColorMapper, Slider
from bokeh.plotting import column, curdoc, figure, row
import colorcet
from dataclasses import dataclass
import datashader as ds
import numpy as np
import pandas as pd
hows = ["linear", "log", "eq_hist", "eq_hist_new"]
@dataclass(frozen=True)
class Count:
canvas_size: int
agg: np.ndarray
mask: np.ndarray
unique_values: np.ndarray
unique_counts: np.ndarray
cdf: np.ndarray
@dataclass
class Transformed:
how: str
counts: np.ndarray
agg: np.ndarray
unique_values: np.ndarray
def create_data(canvas_size):
nsamples = 10000
rng = np.random.default_rng(1289242)
dists = {cat: pd.DataFrame(dict([('x', rng.normal(x, s, nsamples)),
('y', rng.normal(y, s, nsamples)),
('val', val),
('cat', cat)]))
for x, y, s, val, cat in
[( 2, 2, 0.03, 10, "d1"),
( 2, -2, 0.10, 20, "d2"),
( -2, -2, 0.50, 30, "d3"),
( -2, 2, 1.00, 40, "d4"),
( 0, 0, 3.00, 50, "d5")] }
df = pd.concat(dists, ignore_index=True)
df["cat"]=df["cat"].astype("category")
cvs = ds.Canvas(plot_width=canvas_size, plot_height=canvas_size)
agg = cvs.points(source=df, x="x", y="y")
agg = agg.data.ravel()
mask = agg != 0 # Mask of non-zero values
agg = agg[mask] # Only non-zero values
unique_values, unique_counts = np.unique(agg, return_counts=True) # Unique values
cdf = np.cumsum(unique_counts) / unique_counts.sum()
return Count(canvas_size, agg, mask, unique_values, unique_counts, cdf)
def process(how, count, eq_hist_new_power):
# eq_hist_new_power only used if how == "eq_hist_new".
counts = count.unique_counts
if how == "linear":
transformed_unique_values = count.unique_values.copy()
transformed_agg = count.agg.copy()
elif how == "log":
transformed_unique_values = np.log1p(count.unique_values)
transformed_agg = np.log1p(count.agg)
elif how == "eq_hist":
# Transform based on CDF
cdf = count.cdf
cdf = (cdf - cdf[0]) / (cdf[-1] - cdf[0]) # Normalised to range 0..1
transformed_unique_values = np.interp(count.unique_values, count.unique_values, cdf)
transformed_agg = np.interp(count.agg, count.unique_values, cdf)
elif how == "eq_hist_new":
# Transform based on CDF
counts = counts**eq_hist_new_power # Raise to power <= 1
cdf = np.cumsum(counts)
cdf = (cdf - cdf[0]) / (cdf[-1] - cdf[0]) # Normalised to range 0..1
transformed_unique_values = np.interp(count.unique_values, count.unique_values, cdf)
transformed_agg = np.interp(count.agg, count.unique_values, cdf)
else:
raise RuntimeError("Not implemented")
return Transformed(how, counts, transformed_agg, transformed_unique_values)
def shade(count, transformed, cmap):
low = transformed.agg.min()
high = transformed.agg.max()
rspan, gspan, bspan = np.array(list(zip(*map(ds.colors.rgb, cmap))))
span = np.linspace(low, high, len(cmap))
r = np.interp(transformed.agg, span, rspan, left=255).astype(np.uint8)
g = np.interp(transformed.agg, span, gspan, left=255).astype(np.uint8)
b = np.interp(transformed.agg, span, bspan, left=255).astype(np.uint8)
a = np.full_like(r, 255)
rgba = np.column_stack([r, g, b, a])
canvas_size = count.canvas_size
image = np.zeros((canvas_size*canvas_size), dtype=np.uint32)
view = image.view(dtype=np.uint8).reshape((canvas_size*canvas_size, 4))
for i in range(4):
view[count.mask, i] = rgba[:, i]
image.shape = (canvas_size, canvas_size)
return image
count = None
data_cds = {}
image_cds = {}
histogram_cds = {}
data_color_mappers = {}
histogram_color_mappers = {}
def create_all_data():
global count
count = create_data(canvas_size)
for how in hows:
transformed = process(how, count, eq_hist_new_power)
image = shade(count, transformed, cmap)
hist, edges = np.histogram(transformed.agg, bins=nbins)
if how == "eq_hist_new":
hist = hist**eq_hist_new_power
data = dict(count=count.unique_values, transformed=transformed.unique_values,
cdf=count.cdf)
if how in data_cds:
data_cds[how].data = data
else:
data_cds[how] = ColumnDataSource(data)
data = dict(image=[image])
if how in image_cds:
image_cds[how].data = data
else:
image_cds[how] = ColumnDataSource(data)
data = dict(hist=hist, left=edges[:-1], right=edges[1:])
if how in histogram_cds:
histogram_cds[how].data = data
else:
histogram_cds[how] = ColumnDataSource(data)
def create_color_mappers():
for how in hows:
low = data_cds[how].data["transformed"][0]
high = data_cds[how].data["transformed"][-1]
if how in data_color_mappers:
data_color_mappers[how].low = low
data_color_mappers[how].high = high
else:
data_color_mappers[how] = LinearColorMapper(palette=cmap, low=low, high=high)
low = histogram_cds[how].data["left"][0]
high = histogram_cds[how].data["left"][-1]
if how in histogram_color_mappers:
histogram_color_mappers[how].low = low
histogram_color_mappers[how].high = high
else:
histogram_color_mappers[how] = LinearColorMapper(palette=cmap, low=low, high=high)
nbins = 100
canvas_size = 200
cmap = colorcet.rainbow
eq_hist_new_power = 0.5
h = 200
w = 350
lw = 2 # line width
ms = 7 # marker size
lc = "silver" # line color
create_all_data()
create_color_mappers()
cols = []
for i, how in enumerate(hows):
# Datashaded image.
p0 = figure(width=w, height=w, title=f"how = '{how}'", toolbar_location=None)
p0.x_range.range_padding = p0.y_range.range_padding = 0
p0.image_rgba(source=image_cds[how], image="image", x=0, y=0, dw=canvas_size, dh=canvas_size)
# Transformed data space.
p1 = figure(width=w, height=h, x_axis_label="Data space", y_axis_label="Transformed data space",
title="Transform", toolbar_location=None)
p1.line(source=data_cds[how], x="count", y="transformed", line_width=lw, color=lc)
p1.scatter(source=data_cds[how], x="count", y="transformed", marker="o", size=ms,
color=dict(field="transformed", transform=data_color_mappers[how]))
# Linear histogram.
p2 = figure(width=w, height=h, title="Linear histogram in transformed data space",
x_axis_label="Transformed data space (= color space)", y_axis_label="Pixel counts",
toolbar_location=None)
p2.quad(source=histogram_cds[how], top="hist", bottom=0, left="left", right="right",
color=dict(field="left", transform=histogram_color_mappers[how]))
# Log histogram.
p3 = figure(width=w, height=h, title="Log histogram in transformed data space",
x_axis_label="Transformed data space (= color space)", y_axis_label="Pixel counts",
y_axis_type="log", toolbar_location=None)
p3.quad(source=histogram_cds[how], top="hist", bottom=1, left="left", right="right",
color=dict(field="left", transform=histogram_color_mappers[how]))
# CDF.
p4 = figure(width=w, height=h, title="CDF in transformed data space", y_axis_label="CDF",
x_axis_label="Transformed data space (= color space)", toolbar_location=None)
p4.line(source=data_cds[how], x="transformed", y="cdf", line_width=lw, color=lc)
p4.step(source=data_cds[how], x="transformed", y="cdf", line_width=lw, mode="after", color=lc)
p4.scatter(source=data_cds[how], x="transformed", y="cdf", marker="o", size=ms,
color=dict(field="transformed", transform=data_color_mappers[how]))
col = column([p0, p1, p2, p3, p4])
cols.append(col)
def power_callback(_attr, _old, new_value):
global eq_hist_new_power
eq_hist_new_power = new_value
how = "eq_hist_new"
transformed = process(how, count, eq_hist_new_power)
data_cds[how].data = dict(
count=count.unique_values, transformed=transformed.unique_values, cdf=count.cdf)
image = shade(count, transformed, cmap)
image_cds[how].data = dict(image=[image])
hist, edges = np.histogram(transformed.agg, bins=nbins)
if how == "eq_hist_new":
hist = hist**eq_hist_new_power
histogram_cds[how].data = dict(hist=hist, left=edges[:-1], right=edges[1:])
def canvas_size_callback(_attr, _old, new_value):
global canvas_size
canvas_size = new_value
create_all_data()
create_color_mappers()
# Colorbar column at end, just using an image.
color_mapper = LinearColorMapper(palette=cmap, low=0, high=1)
image = np.linspace(0.0, 1.0, len(cmap))
image = np.expand_dims(image, axis=0)
p = figure(width=w, height=80, toolbar_location=None, x_range=(0, 1), y_range=(0, 1),
title="Colormap")
p.yaxis.visible = False
p.grid.visible = False
p.image(image=[image], color_mapper=color_mapper, x=0, y=0, dw=1, dh=1)
power_slider = Slider(start=0.0, end=1.0, value=eq_hist_new_power, step=0.01,
title="eq_hist_new power")
power_slider.on_change("value", power_callback)
canvas_size_slider = Slider(start=50, end=500, value=canvas_size, step=50, title="Canvas size")
canvas_size_slider.on_change("value", canvas_size_callback)
div = Div(text="<p>eq_hist_new raises the histogram bin counts to a power in the range "
"0 <= power <= 1. This reduces the height of the larger counts with respect to the lower "
"counts, bringing the histogram bins closer together where there were large gaps.</p><p>"
"power=1 corresponds to the original eq_hist algorithm, power=0 makes all counts the same, "
"power=0.5 takes square root of the counts.</p>", width=300)
col = column([p, power_slider, canvas_size_slider, div])
cols.append(col)
curdoc().add_root(row(cols))