INNER CODE UNIT · Python

translate_jvp

HIPS/autograd · autograd/core.py:166

def translate_jvp(jvpfun, fun, argnum):
    if jvpfun is None:
        return lambda g, ans, *a, **k: vspace(ans).zeros()
    elif jvpfun == "same":
        return lambda g, ans, *args, **kwargs: fun(*subval(args, argnum, g), **kwargs)
    elif callable(jvpfun):
        return jvpfun
    else:
        raise TypeError(f"Bad JVP '{jvpfun}' for '{fun.__name__}'")


def def_linear(fun):
    """Flags that a function is linear wrt all args"""
    defjvp_argnum(fun, lambda argnum, g, ans, args, kwargs: fun(*subval(args, argnum, g), **kwargs))


# -------------------- vector behavior --------------------

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…