Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion arkouda-env-dev.yml
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ dependencies:
- versioneer
- matplotlib>=3.3.2
- h5py>=3.7.0
- hdf5>=1.12.2
- hdf5==1.14.6
- pip
- types-tabulate
- pytables>=3.10.0
Expand Down
98 changes: 98 additions & 0 deletions arkouda/pandas/groupbyclass.py
Original file line number Diff line number Diff line change
Expand Up @@ -268,6 +268,7 @@ class GroupByReductionType(enum.Enum):
FIRST = "first"
MODE = "mode"
UNIQUE = "unique"
MIN_MEAN_MAX = "min_mean_max"

def __str__(self) -> str:
"""
Expand Down Expand Up @@ -1305,6 +1306,103 @@ def max(self, values: pdarray, skipna: bool = True) -> Tuple[groupable, pdarray]
k, v = self.aggregate(values, "max", skipna)
return k, cast(pdarray, v)

def min_mean_max(self, values: pdarray, skipna: bool = True) -> Tuple[groupable, pdarray, pdarray, pdarray]:
"""
Group another array of values and compute min, mean, and max for each group in one server pass.

Group using the permutation stored in the GroupBy instance.

Parameters
----------
values : pdarray
The values to group and compute stats
skipna: bool
boolean which determines if NANs should be skipped

Returns
-------
Tuple[groupable, pdarray, pdarray, pdarray]
unique_keys : (list of) pdarray or Strings
The unique keys, in grouped order
group_mins : pdarray
One minimum per unique key in the GroupBy instance
group_means : pdarray, float64
One mean value per unique key in the GroupBy instance
group_maxs : pdarray
One maximum per unique key in the GroupBy instance

Raises
------
TypeError
Raised if the values array is not a pdarray object
ValueError
Raised if the key array size does not match the values size
RuntimeError
Raised if min_mean_max is not supported for the values dtype

Notes
-----
This function computes min, mean, and max in a single server-side pass,
which is more efficient than calling min(), mean(), and max() separately.
The mean is always returned as float64 dtype.

Examples
--------
>>> import arkouda as ak
>>> a = ak.randint(1, 5, 10, seed=1)
>>> a
array([2 4 4 2 1 4 1 2 4 3])
>>> g = ak.GroupBy(a)
>>> b = ak.randint(1, 10, 10, seed=1)
>>> b
array([5 7 7 5 2 7 2 5 7 6])
>>> keys, mins, means, maxs = g.min_mean_max(b)
>>> keys
array([1 2 3 4])
>>> mins
array([2 5 6 7])
>>> means
array([2.00000000000000000 5.00000000000000000 6.00000000000000000 7.00000000000000000])
>>> maxs
array([2 5 6 7])

"""
from arkouda.core.client import generic_msg

if values.dtype == bool:
raise TypeError("min_mean_max is only supported for pdarrays of dtype float64, uint64, and int64")

if cast(pdarray, values).size != self.length:
raise ValueError("Attempt to group array using key array of different length")

if self.assume_sorted:
permuted_values = cast(pdarray, values)
else:
permuted_values = cast(pdarray, values)[cast(pdarray, self.permutation)]

rep_msg = generic_msg(
cmd="segmentedReduction",
args={
"values": permuted_values,
"segments": self.segments,
"op": "min_mean_max",
"skip_nan": skipna,
"ddof": 1,
},
)
self.logger.debug(rep_msg)

# Parse the response: should be "min_name+mean_name+max_name"
parts = cast(str, rep_msg).split("+")
if len(parts) != 3:
raise RuntimeError(f"Unexpected response from min_mean_max reduction: {rep_msg}")

mins = create_pdarray(cast(str, parts[0]))
means = create_pdarray(cast(str, parts[1]))
maxs = create_pdarray(cast(str, parts[2]))

return self.unique_keys, mins, means, maxs

def argmin(self, values: pdarray) -> Tuple[groupable, pdarray]:
"""
Group another array of values and return the location of the first minimum of each group.
Expand Down
186 changes: 186 additions & 0 deletions src/ReductionMsg.chpl
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,74 @@ module ReductionMsg
return ret;
}
}
/* segMinMeanMax: Compute min, mean, and max of each segment.
Returns a tuple of three distributed arrays (mins, means, maxs) of the same size as segments.
*/
proc segMinMeanMax(ref values:[] ?t, segments:[?D] int, skipNan=false): ([D] t, [D] real, [D] t) throws {
var mins = makeDistArray(D, t);
var means = makeDistArray(D, real);
var maxs = makeDistArray(D, t);

if (D.size == 0) {
return (mins, means, maxs);
}

// Sentinels are the output for empty/all-NaN segments.
if isRealType(t) {
forall i in D {
mins[i] = +nan:t;
maxs[i] = -nan:t;
}
} else {
forall i in D {
mins[i] = max(t);
maxs[i] = min(t);
}
}

// Each element carries:
// (resetAtSegmentStart, hasValid, min, sum, max, count)
var flagvalues = makeDistArray(values.domain, (bool, bool, t, real, t, int));
forall (fv, v) in zip(flagvalues, values) {
if isRealType(t) && skipNan && isNan(v) {
fv = (false, false, 0:t, 0.0, 0:t, 0);
} else {
fv = (false, true, v, v:real, v, 1);
}
}

forall s in segments with (var agg = newDstAggregator(bool)) {
agg.copy(flagvalues[s][0], true);
}

const scanresult = ResettingMinMeanMaxScanOp scan flagvalues;

forall (i, mn, mu, mx, low) in zip(D, mins, means, maxs, segments)
with (var minAgg = newSrcAggregator(t),
var meanAgg = newDstAggregator(real),
var maxAgg = newSrcAggregator(t)) {
var vi: int;
if (i < D.high) {
vi = segments[i+1] - 1;
} else {
vi = values.domain.high;
}

if (vi >= low) {
const stats = scanresult[vi];
const hasValid = stats(1);
if hasValid {
minAgg.copy(mn, stats(2));
maxAgg.copy(mx, stats(4));
meanAgg.copy(mu, stats(3) / stats(5):real);
} else {
meanAgg.copy(mu, 0.0);
}
}
}

return (mins, means, maxs);
}

@arkouda.registerCommand
proc prodAll(const ref x:[?d] ?t, skipNan: bool): reductionReturnType(t) throws
Expand Down Expand Up @@ -571,6 +639,20 @@ module ReductionMsg
var res = segTail(values.a, segments.a, n);
st.addEntry(rname, createSymEntry(res));
}
when "min_mean_max" {
var (mins, means, maxs) = segMinMeanMax(values.a, segments.a);
var min_name = st.nextName();
var mean_name = st.nextName();
var max_name = st.nextName();
st.addEntry(min_name, createSymEntry(mins));
st.addEntry(mean_name, createSymEntry(means));
st.addEntry(max_name, createSymEntry(maxs));
var repMsg = "created " + st.attrib(min_name)
+ "+created " + st.attrib(mean_name)
+ "+created " + st.attrib(max_name);
rmLogger.debug(getModuleName(),getRoutineName(),getLineNumber(),repMsg);
return new MsgTuple(repMsg, MsgType.NORMAL);
}
otherwise {
var errorMsg = notImplementedError(pn,op,gVal.dtype);
rmLogger.error(getModuleName(),getRoutineName(),getLineNumber(),errorMsg);
Expand Down Expand Up @@ -641,6 +723,20 @@ module ReductionMsg
var res = segCount(segments.a, values.size);
st.addEntry(rname, createSymEntry(res));
}
when "min_mean_max" {
var (mins, means, maxs) = segMinMeanMax(values.a, segments.a);
var min_name = st.nextName();
var mean_name = st.nextName();
var max_name = st.nextName();
st.addEntry(min_name, createSymEntry(mins));
st.addEntry(mean_name, createSymEntry(means));
st.addEntry(max_name, createSymEntry(maxs));
var repMsg = "created " + st.attrib(min_name)
+ "+created " + st.attrib(mean_name)
+ "+created " + st.attrib(max_name);
rmLogger.debug(getModuleName(),getRoutineName(),getLineNumber(),repMsg);
return new MsgTuple(repMsg, MsgType.NORMAL);
}
otherwise {
var errorMsg = notImplementedError(pn,op,gVal.dtype);
rmLogger.error(getModuleName(),getRoutineName(),getLineNumber(),errorMsg);
Expand Down Expand Up @@ -695,6 +791,20 @@ module ReductionMsg
var res = segCount(segments.a, values.size) - nanCounts(values.a, segments.a);
st.addEntry(rname, createSymEntry(res));
}
when "min_mean_max" {
var (mins, means, maxs) = segMinMeanMax(values.a, segments.a, skipNan);
var min_name = st.nextName();
var mean_name = st.nextName();
var max_name = st.nextName();
st.addEntry(min_name, createSymEntry(mins));
st.addEntry(mean_name, createSymEntry(means));
st.addEntry(max_name, createSymEntry(maxs));
var repMsg = "created " + st.attrib(min_name)
+ "+created " + st.attrib(mean_name)
+ "+created " + st.attrib(max_name);
rmLogger.debug(getModuleName(),getRoutineName(),getLineNumber(),repMsg);
return new MsgTuple(repMsg, MsgType.NORMAL);
}
otherwise {
var errorMsg = notImplementedError(pn,op,gVal.dtype);
rmLogger.error(getModuleName(),getRoutineName(),getLineNumber(),errorMsg);
Expand Down Expand Up @@ -989,6 +1099,82 @@ module ReductionMsg
}
}

/* Performs a segmented scan where each element tracks
* (hasValid, min, sum, max, count) and segment boundaries reset state.
*/
class ResettingMinMeanMaxScanOp: ReduceScanOp {
type eltType;
var value: eltType;

proc identity {
return (false, false, 0:eltType(2), 0.0, 0:eltType(4), 0);
}

proc combineStats(hasValidA, minA, sumA, maxA, countA,
hasValidB, minB, sumB, maxB, countB) {
if hasValidA {
if hasValidB {
return (true, min(minA, minB), sumA + sumB, max(maxA, maxB), countA + countB);
} else {
return (true, minA, sumA, maxA, countA);
}
} else {
if hasValidB {
return (true, minB, sumB, maxB, countB);
} else {
return (false, minA, sumA, maxA, countA);
}
}
}

proc accumulate(x) {
const (resetB, hasValidB, minB, sumB, maxB, countB) = x;
const (hasResetA, hasValidA, minA, sumA, maxA, countA) = value;

if resetB {
value = (hasResetA | resetB, hasValidB, minB, sumB, maxB, countB);
} else {
const (hasValid, mn, sm, mx, ct) = combineStats(hasValidA, minA, sumA, maxA, countA,
hasValidB, minB, sumB, maxB, countB);
value = (hasResetA | resetB, hasValid, mn, sm, mx, ct);
}
}

proc accumulateOntoState(ref state, x) {
const (prevReset, hasValidB, minB, sumB, maxB, countB) = x;
const (hasResetA, hasValidA, minA, sumA, maxA, countA) = state;

if hasResetA {
state = (hasResetA | prevReset, hasValidA, minA, sumA, maxA, countA);
} else {
const (hasValid, mn, sm, mx, ct) = combineStats(hasValidA, minA, sumA, maxA, countA,
hasValidB, minB, sumB, maxB, countB);
state = (hasResetA | prevReset, hasValid, mn, sm, mx, ct);
}
}

proc combine(x) {
const (xHasReset, hasValidB, minB, sumB, maxB, countB) = x.value;
const (hasResetA, hasValidA, minA, sumA, maxA, countA) = value;

if hasResetA {
value = (hasResetA | xHasReset, hasValidA, minA, sumA, maxA, countA);
} else {
const (hasValid, mn, sm, mx, ct) = combineStats(hasValidA, minA, sumA, maxA, countA,
hasValidB, minB, sumB, maxB, countB);
value = (hasResetA | xHasReset, hasValid, mn, sm, mx, ct);
}
}

proc generate() {
return value;
}

proc clone() {
return new unmanaged ResettingMinMeanMaxScanOp(eltType=eltType);
}
}

proc segProduct(values:[] ?t, segments:[?D] int, skipNan=false): [D] real throws {
/* Compute the product of values in each segment. The logic here
is to convert the product into a sum in the log-domain. To
Expand Down