"""适合面试现场书写和测试的 NumPy 小函数。"""
import numpy as np


def softmax(x, axis=-1):
    x=np.asarray(x,dtype=np.float64)
    shifted=x-np.max(x,axis=axis,keepdims=True)
    exp=np.exp(shifted)
    return exp/exp.sum(axis=axis,keepdims=True)


def top_k_indices(x, k):
    x=np.asarray(x)
    if not 1<=k<=x.size: raise ValueError("k 越界")
    ids=np.argpartition(-x,k-1)[:k]
    return ids[np.argsort(-x[ids])]


if __name__ == "__main__":
    p=softmax([[1000,1001,999]])
    assert np.all(np.isfinite(p)) and np.allclose(p.sum(-1),1)
    assert list(top_k_indices([1,9,3,7],2))==[1,3]
    print("interview drill tests passed")

