On Monads, Monoids and Endofunctors 3: The monad

10 minute read

Published:

Spoiler: Category theory has applications in machine learning engineering

Note: The code lives in ianqs - applied_category_theory/3_monads

Recap

You’re an engineer on an ML team working on distributed hyperparameter search at #LARGE_CORPORATION, and we want to make sure our code is web scale.

In the monoid post we built a Summary that folds validation losses up a reduction tree, so that the individual worker machines and the accumulators (reducers) can all use the same code. In the followup functor post we swapped the container from a list to a LossTree without breaking the map/filter contract; our downstream consumers never had to know or care what exactly our code “did” and neither did we, so long as everyone follows the contract.

One huge simplification in the monoid post was that we only cared about the statistics, not WHY a run might have failed: an empty Summary object was enough. (If that wasn’t a problem setup, I don’t know what is)

# From the monoid post
def run_config(config: dict) -> Summary:
    if random.random() < 0.1:
        return Summary()  # a failure just becomes the identity
    val = random.random()
    return Summary.of(loss=val)

New Requirements

Because you’re such a go-getter and you keep volunteering for more work, the bosses tasked you with figuring out a clean way to store WHY our runs have failed. Was there an OOM, a code crash, or the datacenter down the road overloaded the entire grid? And since you built the rest of the architecture, it really only makes sense that you finish it off and hopefully get that juicy promotion.

Setup

At the very end of the monoid post we outlined a few steps in the pipeline: load_shard, train, validate and score, so let’s just go with those for now - the reasoning and core principles are the same, anyways. Each of the aforementioned steps comes with their own potential failure modes:

def load_shard(config): ...   # might fail on "shard missing"
def train(run): ...           # might fail on "diverged (NaN loss)"
def validate(run): ...        # might fail on "OOM" (I guess? This is just pedagogical anyways)
def score(run): ...           # assigns the final loss; the one step that can't fail

Solution: The Naive Method

The standard way to handle the failure is to create a “ladder” of the results; each of the aforementioned functions can be made to return a tuple of two elements:

MaybeResult = Run | None   # the run if the step succeeded, or None if it failed
MaybeFailure = str | None  # why it failed, or None if it didn't

so each step hands back the result in the first slot and the reason in the second, and exactly one of the two is ever a None.

# Assume this is the signature for the 4 functions above:

def some_func(fn_in) -> tuple[MaybeResult, MaybeFailure]:
    ...
def evaluate(config):
    # Alternatively, we can check "reason" instead of the computation result
    shard, reason = load_shard(config)
    if shard is None:
        return None, reason
    trained, reason = train(shard)
    if trained is None:
        return None, reason
    validated, reason = validate(trained)
    if validated is None:
        return None, reason
    return score(validated)

This is … fine? It’s just a little … gauche. We know we can do better.

The Result

What do we want, exactly? It’d be nice if we had all the boring “plumbing” work pushed into one place, so that a failure gets an actual type instead of a None and some arbitrary str. We can take inspiration from Rust’s Result, which is either an Ok carrying a value, or an Err carrying a reason.

@dataclass(frozen=True)
class Ok:
    value: object

    def map(self, f) -> Result:
        return Ok(f(self.value))


@dataclass(frozen=True)
class Err:
    reason: str

    def map(self, f) -> Err:
        return self   # a dead run stays dead, and keeps its reason


Result = Ok | Err

If the structure looks familiar implementation-wise, it’s because Result is a functor! Ok and Err are its two cases, and we don’t have to care about the underlying implementation of either of them - so long as Result has a map we’re good to go. Remember how in the previous post on functors they were lists/iterables/BSTs, but we didn’t need to care as long as they provided a map that tells us how to apply functions onto the contents? Same principles here.

Why the functor wasn’t enough

Cool cool cool. So it looks like a functor is all we need, then? Well, not quite. Consider the following example:

load_shard(config).map(train)

# Option 1: Ok(Ok(run))
# Option 2: Ok(Err("diverged (NaN loss)"))

What does that output? It’s a Result inside a Result, and our map was never actually made to support this operation. In fact, if we took a look at the function signature, this would already have been clear. A natural response might be to modify map to unwrap the value, but then that breaks our OTHER functions that rely on map i.e. it violates our “contract”.

This is where the monad comes in and it does so by defining a flatten (typically called join , but flatten is just such an intuitive name for it)

def flatten(nested: Result) -> Result:
    if isinstance(nested, Err):
        return nested
    return nested.value        # Ok(inner) -> inner

result.py opens with the two ways people usually introduce a monad: the direct definition, which hands you unit and bind as primitives, and the functor route, which hands you unit and join (our flatten) on top of map. With map and flatten in hand, the direct definition’s bind is just the composition of the two, and it’s the same line in both classes:

    def bind(self, step):
        return flatten(self.map(step))

On an Ok, map runs the step and leaves its Result double-wrapped, then flatten undoes the extra layer, so bind amounts to step(self.value). On an Err, map does nothing and flatten passes the Err straight through, so the first failure slides to the end of the pipeline still carrying its reason. The coolest part (to me) is that we didn’t write explicit if-else skips in the pipeline itself: we just used Result and let the types define how code should interact with their wrapped data (via the map and flatten). Each step returns a valid Result and our old if-else skips flatten into a simple direct chain:

def evaluate(config) -> Result:
    return (load_shard(config)
            .bind(train)
            .bind(validate)
            .bind(score))

Remember how earlier we asked the question “So it looks like a functor is all we need, then?” We got to it by showing that a functor had almost everything we needed, except it needed a flatten (aka join). Well, we can just think of a monad as a functor with a flatten operation.

The Haskell signature

Like in the functor post, we should look at the function signature (this time, of bind):

(>>=) :: m a -> (a -> m b) -> m b

From left-to-right, bind involves 1) an m a, a container m holding an a (our Ok holding a value); 2) a function a -> m b that transforms an unwrapped value into a new value in the same container type; 3) a non-nested b in the same container.

Heck, if we put it next to the fmap from our functor post: fmap :: (a -> b) -> f a -> f b, we see that the only difference is that the function returns m b instead of a plain b. Okay, the container and the function swap places in the signature, but that’s cosmetic. And I suppose there’s also a unit, but that’s just a trivial constructor.

A monoid in the category of endofunctors

Remember the meme the first post left you with?

A monad is a monoid in the category of endofunctors, what’s the problem?

Four years after my first post we’re about ready to tie it up. Look at the two things we added to turn our functor into a monad: Ok lifts a plain value into a Result, and flatten collapses two nested layers into one. An identity, and an associative combine. Those are the monoid laws, which is why these assertions hold:

def test_monad_laws():
    def grow(n): return Ok(n + 1)
    def double(n): return Ok(n * 2)

    assert Ok(3).bind(grow) == grow(3)                      # left identity
    assert Ok(3).bind(Ok) == Ok(3)                          # right identity
    assert (Ok(3).bind(grow).bind(double)              # associativity
            == Ok(3).bind(lambda n: grow(n).bind(double)))

It’s the same diagram we drew in the monoid post, one level of abstraction up:

  • Back then: M was Summary, $\bigotimes$ put two of them side by side, and $\mu$ combined them. There, the identity is an empty Summary (unit)

  • Now M is the functor Result, $\bigotimes$ nests one Result inside another, and $\mu$ is flatten (join), and the identity here is Ok (unit).

Note admittedly, the diagram doesn’t actually show the left-identity and right-identity, but that’s arguably not important for this discussion

monoid diagram

Retrospective

The monoid post ended with

I [sic] would be nice to be able to just hand off the crash and have the later steps know what to do with it; no more if None … else.

and arguably we’ve done exactly this:

def laddered_evaluate(config):
    # Alternatively, we can check "reason" instead of the computation result
    shard, reason = load_shard(config)
    if shard is None:
        return None, reason
    trained, reason = train(shard)
    if trained is None:
        return None, reason
    validated, reason = validate(trained)
    if validated is None:
        return None, reason
    return score(validated)

def monadic_evaluate(config) -> Result:
    return (load_shard(config)
            .bind(train)
            .bind(validate)
            .bind(score))

However, there’s just so much more to cover! We can talk about combning of these category theory structures for our problem space, and about the different KINDS of monads (we focused on the kind of monad that lets us avoid the if None … else issue, but monads can also carry state, model side-effects and much more). In the next post (which I didn’t even intend to write, but this exposition is already wildly long), we bring the band back together and start free-styling these structures: the Summary monoid, the LossTree functor and the Result monad fold into a single Report, and we meet a second monad, the Writer. Stay tuned!