~/searching/sorting-keys

Sorting with custom keys

Describe the order and let the library sort. Keys, tuples, stable sorts, comparators, and huge values squeezed into ranks 0..k−1.

what

Give the built-in sort a key, or a rule for two items, instead of writing a sort. Coordinate compression swaps each value for its rank among the distinct values.

use when

Items must be ordered by several fields or by a computed value, or huge values need to become small array indexes.

time

O(n log n)

space

O(n)

You’ll recognise it when

  • The output must be in an order set by several fields: most points first, then names from A to Z.
  • The order depends on a computed value: distance from a point, length, the number hidden inside a file name.
  • You must arrange items so that the combined result is as big or as small as possible, and the only rule you have says which of two items should go first.
  • Values are huge or sparse, up to 10¹⁸, but only their order matters, and you want to use them as array indexes.
  • Once the items are in order, the rest is easy: equal items sit together, and a single greedy or two-pointer pass finishes the job.

If you only need the k smallest items, sorting everything is more work than needed. A heap of size k does it in O(n log k).

The idea

A teacher sorting exam papers doesn’t invent a sorting method. She decides what matters: class first, then surname within each class. Then she sorts the usual way. You do the same. The library sort is fast and correct, so your job is to describe the order, not to write a sort.

There are two ways to describe it. A key turns each item into a value that already sorts the right way. A comparator looks at two items and says which goes first.

Keys and tuples

A key is computed once per item, and then the keys are compared. Tuples compare field by field: the first fields decide, and only a tie moves on to the next field. So a tuple key sorts by one field, then another. To sort a number field in descending order, negate it.

teams = [("owls", 7), ("bees", 9), ("ants", 7)]
# Points from high to low, then names from A to Z.
teams.sort(key=lambda t: (-t[1], t[0]))
# [('bees', 9), ('ants', 7), ('owls', 7)]
# Nearest the origin first. Squared distance needs no sqrt.
points.sort(key=lambda p: p[0] ** 2 + p[1] ** 2)

A key can be anything that compares: a number, a string, a tuple, even a list. Splitting a file name into text and number parts, with the numbers as int, puts img9 before img10.

sorted(a) returns a new list and leaves a alone. a.sort() sorts in place and returns None, so a = a.sort() loses your list.

The same order in C++ and Java:

// More points first, then names from A to Z.
sort(teams.begin(), teams.end(), [](auto& a, auto& b) {
if (a.points != b.points) return a.points > b.points;
return a.name < b.name;
});
teams.sort(Comparator.comparingInt((Team t) -> t.points).reversed()
.thenComparing(t -> t.name));

In Java, .reversed() flips everything chained before it, so put it straight after the key it should flip. The (Team t) type is needed there: without it, Java can’t work out the lambda’s type before .reversed().

Stable sorts

A sort is stable when items with equal keys keep their original order. That matters in two ways:

  • Ties keep input order. Sort submissions by score, and among equal scores the earlier submission stays first, with no extra field.
  • Several passes build a multi-key sort. Sort by the least important field first, then by the most important. The second sort keeps the first sort’s order inside each tie. This is how you get “points high to low, then name A to Z” when you can’t negate a string.
Stable?
Python sorted, list.sort Always, even with reverse=True
C++ std::sort No. Use std::stable_sort when ties must keep their order
Java List.sort, Collections.sort, Arrays.sort on objects Yes
Java Arrays.sort on int[], long[] No, but equal numbers can’t be told apart anyway

When no key works: comparators

Arrange the numbers 3, 30 and 34 so that writing them side by side gives the biggest number. Sorting their strings from high to low gives 34, 30, 3, which spells 34303. But 34330 is bigger: 3 should come before 30. No simple per-item key captures that, but a rule for two items does: put a before b when a + b > b + a as strings. "3" + "30" is "330", which beats "303", so 3 goes first.

from functools import cmp_to_key
# Negative: a goes first. Positive: b goes first. 0: a tie.
def order(a, b):
if a + b > b + a:
return -1
if a + b < b + a:
return 1
return 0
strs.sort(key=cmp_to_key(order))

Python 3’s sort only takes a key, so functools.cmp_to_key wraps the two-item rule as one. Prefer a plain key= when one exists: a key is computed once per item, while a comparator runs on every comparison, about n log n times.

A comparator must describe a real order. If a goes before b and b before c, then a must go before c (it’s transitive), and ties must be transitive too. The sort only compares some pairs and trusts the rest. Break the rule and the output depends on which pairs it happened to compare. Java may throw “Comparison method violates its general contract!”, and in C++ it’s undefined behaviour. The concatenation rule is safe: it can be proved transitive. “Treat values within 0.5 of each other as equal” is not: 1.0 ties with 1.4 and 1.4 ties with 1.8, but 1.0 doesn’t tie with 1.8.

Python C++ Java
By a key key=f f(a) < f(b) in a lambda Comparator.comparing(f)
Then another key a tuple key compare field by field .thenComparing(g)
Descending negate the key, or reverse=True > instead of < .reversed()
Rule for two items cmp_to_key, returns negative, 0 or positive returns true if a goes strictly first returns negative, 0 or positive

How it works

We’ll write coordinate compression: replace every value with its rank among the distinct values. Values as big as 10¹⁸ become the small numbers 0 to k − 1, and the order is unchanged, so you can use them as array indexes.

Scroll through the steps and the graphic follows along. You can also press play, or edit the values: try all-equal values, or values already in order.

  1. Copy the input and sort the copy: vals = sorted(xs). Sorting only cares about order, not about how far apart the values are. Equal values now sit side by side.
  2. Walk along vals with v. If v differs from the last value in uniq, it’s new: append it. The first value is always new.
  3. If v equals uniq[-1], it’s a repeat: skip it. Every copy of a value sits in one run, so checking the last kept value is enough.
  4. uniq now holds the k distinct values in increasing order. Go back to the original order and take each x from xs.
  5. Binary search finds x in uniq. Its index there is its rank: the number of distinct values smaller than x. Store it in ranks.
  6. Every value now has a rank from 0 to k − 1, in the positions of xs. Smaller values get smaller ranks, and equal values share one.
loading sorting-keys…

Why it’s correct: uniq is sorted and has no repeats, so the index of x in it is exactly the count of distinct values below x. That count keeps the order: if a < b, every distinct value below a is also below b, and so is a itself. So a gets the smaller rank.

from bisect import bisect_left
def compress(xs):
"""Rank of each value among the distinct values, from 0."""
vals = sorted(xs)
uniq = []
for v in vals:
if uniq and uniq[-1] == v:
continue
uniq.append(v)
ranks = []
for x in xs:
ranks.append(bisect_left(uniq, x))
return ranks
#include <algorithm>
#include <vector>
using namespace std;
// Rank of each value among the distinct values, from 0.
vector<int> compress(const vector<long long>& xs) {
vector<long long> vals = xs;
sort(vals.begin(), vals.end());
vector<long long> uniq;
for (long long v : vals) {
if (!uniq.empty() && uniq.back() == v) {
continue;
}
uniq.push_back(v);
}
vector<int> ranks;
for (long long x : xs) {
auto it = lower_bound(uniq.begin(), uniq.end(), x);
ranks.push_back(it - uniq.begin());
}
return ranks;
}
import java.util.Arrays;
class Compress {
// Rank of each value among the distinct values, from 0.
static int[] compress(long[] xs) {
long[] vals = xs.clone();
Arrays.sort(vals);
long[] uniq = new long[vals.length];
int k = 0; // distinct values so far
for (long v : vals) {
if (k > 0 && uniq[k - 1] == v) {
continue;
}
uniq[k++] = v;
}
int[] ranks = new int[xs.length];
for (int i = 0; i < xs.length; i++) {
ranks[i] = Arrays.binarySearch(uniq, 0, k, xs[i]);
}
return ranks;
}
}

Each language has a shorter way to build uniq: sorted(set(xs)) in Python, sort then u.erase(unique(u.begin(), u.end()), u.end()) in C++, and Arrays.stream(xs).distinct().sorted().toArray() in Java. The loop above is what they do inside.

Why it’s O(n log n)

A sort that works by comparing items needs O(n log n) comparisons. With a key, each comparison compares two keys, and the key function runs only n times. With a comparator, the comparator itself runs O(n log n) times, so an expensive one, like joining two strings of length L, makes the sort O(n log n · L).

For compression, with n values of which k are distinct:

Step Time
Sort the copy O(n log n)
Build uniq O(n)
One binary search per value O(n log k)
Total O(n log n)

Space is O(n) for vals, uniq and ranks.

Common mistakes

Using <= in a C++ comparator

std::sort needs a strict weak ordering: the comparator says whether a goes strictly before b, so it must return false when they’re equal. With <=, equal items each claim to go first. That’s undefined behaviour: a wrong order, or a crash that reads outside the array when many values are equal.

return p <= q; // ✗ true when p == q: undefined behaviour
return p < q; // ✓ false when p == q

Subtracting in a Java comparator

(a, b) -> a - b looks neat, but a - b overflows once the values are more than about 2.1 × 10⁹ apart: 2,000,000,000 − (−2,000,000,000) doesn’t fit in an int. The sign flips and the order silently breaks. Integer.compare or Comparator.naturalOrder() never overflow.

list.sort((a, b) -> a - b); // ✗ can overflow
list.sort((a, b) -> Integer.compare(a, b)); // ✓ no overflow

Letting reverse=True flip the tie-break

reverse=True reverses the whole key, tie-breaker included. With a key of (points, name) you get names from Z to A among equal points. Negate only the field that should go down.

key=lambda t: (t[1], t[0]), reverse=True # ✗ names Z to A too
key=lambda t: (-t[1], t[0]) # ✓ only points flip

Ranking before removing duplicates

Index positions in the sorted list aren’t ranks when values repeat. The dictionary keeps the last index of each value, so for [5, 5, 7] the ranks start at 1 instead of 0. Remove duplicates first.

{v: i for i, v in enumerate(sorted(xs))} # ✗ 5→1, 7→2
{v: i for i, v in enumerate(sorted(set(xs)))} # ✓ 5→0, 7→1

Variations

  • A dictionary instead of binary search. Build rank = {value: i for i, value in enumerate(sorted(set(xs)))} once, then each lookup is O(1) on average. The sort still makes the whole thing O(n log n).
  • Compressing for a tree. A Fenwick tree or segment tree needs small indexes. Compress first, then index the tree by rank: counting, for each value, how many smaller values came before it is the classic use.
  • Sorting indexes. order = sorted(range(n), key=lambda i: xs[i]) sorts positions instead of values, so you still know where each item came from. Use it when the answer must be reported in the original order.
  • Only the k smallest. A heap of size k finds them in O(n log k). Quickselect finds the k-th smallest in O(n) on average, with everything smaller to its left.
  • Writing the sort yourself. Merge sort splits the list in half, sorts each half, then merges them by repeatedly taking the smaller front item. It’s O(n log n) on every input, and stable if ties take from the left half. Quicksort that always picks the first item as its pivot is O(n²) on input that’s already sorted; pick a random pivot instead.

Check yourself

5 quick questions. Pick an answer to see why it's right or wrong.

  1. 1

    You have 10⁵ values between −10⁹ and 10⁹. For each one you must count how many smaller values came before it, using a Fenwick tree indexed by value. What do you do first?

  2. 2

    What does this print?

    rows = [("ann", 3), ("bob", 1), ("cat", 3), ("dan", 1)]
    by_score = sorted(rows, key=lambda r: r[1], reverse=True)
    print([name for name, _ in by_score])
  3. 3

    What does this print?

    from functools import cmp_to_key
    def order(a, b):
    if a + b > b + a:
    return -1
    if a + b < b + a:
    return 1
    return 0
    s = ["3", "30", "34"]
    print("".join(sorted(s, reverse=True)), "".join(sorted(s, key=cmp_to_key(order))))
  4. 4

    This Java sort works in every small test but puts some values in the wrong order when they range from −2 × 10⁹ to 2 × 10⁹. What’s the bug?

    Integer[] a = load();
    Arrays.sort(a, (p, q) -> p - q);
  5. 5

    A C++ program sorts a million scores with sort(v.begin(), v.end(), [](int a, int b) { return a <= b; });. Most scores are 0. It sometimes crashes. Why?

Practice problems

Solve these right here, in Python, C++ or Java. Tests run as you go.

Further reading

esc