在上一章我们学到,Python 的函数是一等公民——可以作为参数传递、作为返回值返回。在这个基础上,**高阶函数(Higher-Order Function)**就是把函数当成普通值来操作的函数:要么接收函数作为参数,要么返回一个函数,或者两者兼备。

本章我们将学习 Python 中常见的高阶函数模式,包括 mapfilterreducesorted 的 key 函数、Lambda 表达式、偏函数,以及函数组合的技巧。

函数作为参数

这是最直观的高阶函数形式——一个函数接收另一个函数作为参数:

def apply_twice(func, value):
    """对 value 应用两次 func"""
    return func(func(value))

def add_one(x):
    return x + 1

print(apply_twice(add_one, 5))  # 7(5→6→7)

def square(x):
    return x ** 2

print(apply_twice(square, 2))   # 16(2→4→16)

这种模式在 Python 标准库中无处不在。比如 sorted()key 参数:

words = ["python", "Java", "C", "javascript", "Go"]

# 按照字符串长度排序
sorted_by_len = sorted(words, key=len)
print(sorted_by_len)  # ['C', 'Go', 'Java', 'python', 'javascript']

# 自定义排序逻辑
sorted_by_last_char = sorted(words, key=lambda s: s[-1])
print(sorted_by_last_char)  # ['Java', 'C', 'python', 'Go', 'javascript']

Lambda 表达式

Lambda 表达式是创建匿名函数的快捷方式。当需要一个简单的、只使用一次的函数时,可以用 lambda 替代完整的 def 定义:

# 完整的函数定义
def double(x):
    return x * 2

# 等价的 lambda
double_lambda = lambda x: x * 2

print(double(5))          # 10
print(double_lambda(5))   # 10

Lambda 语法

lambda 参数列表: 返回值表达式

关键点:

  • 一行代码,不能包含语句(不能有 if 语句、赋值语句等)
  • 只能写一个表达式,表达式的结果就是返回值
  • 不需要 return 关键字

Lambda 示例

# 多个参数
add = lambda a, b: a + b
print(add(3, 5))  # 8

# 带默认参数
greet = lambda name, greeting="你好": f"{greeting}{name}"
print(greet("小明"))           # 你好,小明
print(greet("John", "Hello"))  # Hello,John

# 条件表达式(三元运算符)
max_val = lambda a, b: a if a > b else b
print(max_val(10, 20))  # 20

# 与 sorted 配合
students = [
    {"name": "张三", "score": 85},
    {"name": "李四", "score": 92},
    {"name": "王五", "score": 78},
]

# 按分数排序
ranked = sorted(students, key=lambda s: s["score"], reverse=True)
for s in ranked:
    print(f"{s['name']}: {s['score']}")
# 李四: 92
# 张三: 85
# 王五: 78

Lambda 的限制

# 以下都是不合法的 lambda

# 不能包含赋值语句
# lambda x: x += 1     # SyntaxError

# 不能包含多个语句
# lambda x: print(x); return x * 2  # SyntaxError

# 不能包含循环或 try/except
# lambda x: for i in x: print(i)    # SyntaxError

如果需要比 lambda 更复杂的逻辑,请使用常规的 def 函数。

代码风格建议: 在需要复杂逻辑时应使用 def,lambda 只适合极其简单的转换逻辑。

map —— 映射

map(function, iterable) 将函数应用到可迭代对象的每个元素上,返回一个迭代器:

numbers = [1, 2, 3, 4, 5]

# 每个元素平方
squared = map(lambda x: x ** 2, numbers)
print(list(squared))  # [1, 4, 9, 16, 25]

# 多个可迭代对象——函数必须接收多个参数
list1 = [1, 2, 3]
list2 = [10, 20, 30]
result = map(lambda a, b: a + b, list1, list2)
print(list(result))  # [11, 22, 33]

要不要用 map?

Python 社区中,列表推导式通常是更 Pythonic 的选择

# map 方式
result_map = list(map(lambda x: x ** 2, range(10)))

# 列表推导式——更清晰
result_comp = [x ** 2 for x in range(10)]

print(result_map == result_comp)  # True

什么时候用 map?

  • 已经有了一个现成的函数(不需要 lambda)
  • 需要惰性求值(map 返回迭代器,按需计算)
# 有现成函数时 map 更简洁
nums = ["1", "2", "3", "4", "5"]
parsed = list(map(int, nums))  # 比 [int(x) for x in nums] 更紧凑
print(parsed)  # [1, 2, 3, 4, 5]

filter —— 过滤

filter(function, iterable) 筛选出函数返回 True 的元素:

numbers = range(1, 21)

# 筛选偶数
evens = filter(lambda x: x % 2 == 0, numbers)
print(list(evens))  # [2, 4, 6, 8, 10, 12, 14, 16, 18, 20]

# 筛选正数
values = [-3, 0, 5, -1, 8, 0, 2]
positives = filter(lambda x: x > 0, values)
print(list(positives))  # [5, 8, 2]

# 使用 None 作为函数——过滤掉 falsy 值
mixed = [0, "hello", "", None, 42, [], [1, 2]]
truthy = filter(None, mixed)
print(list(truthy))  # ['hello', 42, [1, 2]]

同样,列表推导式通常更清晰

# filter 方式
evens_filter = list(filter(lambda x: x % 2 == 0, range(20)))

# 列表推导式
evens_comp = [x for x in range(20) if x % 2 == 0]

print(evens_filter == evens_comp)  # True

reduce —— 归约

reduce 位于 functools 模块中。它对可迭代对象做累积运算:

from functools import reduce

# 计算乘积
product = reduce(lambda a, b: a * b, [1, 2, 3, 4, 5])
print(product)  # 120

# 过程:((((1 * 2) * 3) * 4) * 5)

# 找最大值
max_val = reduce(lambda a, b: a if a > b else b, [3, 7, 2, 9, 1])
print(max_val)  # 9

# 使用初始值(第三个参数)
total = reduce(lambda a, b: a + b, [1, 2, 3], 10)
print(total)  # 16(10 + 1 + 2 + 3)

reduce 的通用性不如专门的函数好——如果功能明确(求和用 sum、求积有对应的 math.prod),优先使用专用函数:

from math import prod

# 不推荐
product_reduce = reduce(lambda a, b: a * b, range(1, 6))

# 推荐——更直观
product_direct = prod(range(1, 6))

print(product_direct)  # 120

sorted 的 key 参数

sorted()key 参数是高阶函数最常用的场景之一。它的值是一个函数,该函数接收一个元素并返回用于排序的"键":

data = [
    ("Alice", 28, "Engineer"),
    ("Bob", 35, "Designer"),
    ("Charlie", 22, "Student"),
]

# 按年龄排序
by_age = sorted(data, key=lambda person: person[1])
print(by_age)
# [('Charlie', 22, 'Student'), ('Alice', 28, 'Engineer'), ('Bob', 35, 'Designer')]

# 按名字长度排序
by_name_len = sorted(data, key=lambda person: len(person[0]))
print(by_name_len)
# [('Bob', 35, 'Designer'), ('Alice', 28, 'Engineer'), ('Charlie', 22, 'Student')]

operator 模块

operator 模块提供了一系列高效地获取元素或属性的函数,与 sorted 配合使用非常优雅:

from operator import itemgetter, attrgetter

# itemgetter —— 从序列/映射中取值
data = [
    ("Alice", 28, "Engineer"),
    ("Bob", 35, "Designer"),
    ("Charlie", 22, "Student"),
]

# 等价于 lambda x: x[1]
sorted(data, key=itemgetter(1))  # 按年龄

# 多级排序——先按年龄,再按名字
sorted(data, key=itemgetter(1, 0))

# 配合对象使用
class Person:
    def __init__(self, name, age):
        self.name = name
        self.age = age
    def __repr__(self):
        return f"Person({self.name}, {self.age})"

people = [Person("Alice", 28), Person("Bob", 22), Person("Charlie", 35)]
sorted(people, key=attrgetter("age"))
# [Person(Bob, 22), Person(Alice, 28), Person(Charlie, 35)]

operator 的其他常用函数:

from operator import add, sub, mul, truediv, eq, neg

print(add(5, 3))      # 8
print(mul(4, 7))       # 28
print(neg(10))         # -10

# 与 reduce 结合——用 operator 替代 lambda
from functools import reduce

total = reduce(add, [1, 2, 3, 4, 5])  # 比 reduce(lambda a,b: a+b, ...) 更好
print(total)  # 15

函数作为返回值

高阶函数的另一种形式——返回一个函数。这实际上就是我们在上一章学到的闭包

def make_power(exponent):
    """返回一个计算 exponent 次幂的函数"""
    def power(base):
        return base ** exponent
    return power

square = make_power(2)
cube = make_power(3)

print(square(5))  # 25
print(cube(3))    # 27

functools.partial —— 偏函数

偏函数(Partial Function)是"冻结"一个函数的某些参数,生成一个参数更少的新函数:

from functools import partial

# 原始函数
def power(base, exponent):
    return base ** exponent

# 冻结 exponent=2 和 exponent=3
square = partial(power, exponent=2)
cube = partial(power, exponent=3)

print(square(5))  # 25——等价于 power(5, exponent=2)
print(cube(3))    # 27——等价于 power(3, exponent=3)

实用场景

# 场景:数据库查询需要重复设置连接参数
def connect_db(host, port, user, password, database):
    print(f"连接 {user}@{host}:{port}/{database}")
    # 实际连接逻辑...

# 创建一系列有默认连接的偏函数
local_conn = partial(connect_db, host="localhost", port=5432)
prod_conn = partial(connect_db, host="prod.example.com", port=5432, user="admin")

local_conn(user="dev", password="dev123", database="test")
# 连接 dev@localhost:5432/test

prod_conn(password="secret", database="production")
# 连接 admin@prod.example.com:5432/production

偏函数 vs 默认参数:

偏函数适合在不修改原函数定义的情况下,为函数添加默认值——特别适合适配第三方库。

函数组合(Composition)

函数组合是指将多个函数串联起来,一个函数的输出是另一个函数的输入:

# 手动组合
def compose(f, g):
    """返回 f(g(x))"""
    def composed(x):
        return f(g(x))
    return composed

def double(x):
    return x * 2

def increment(x):
    return x + 1

# double(increment(3)) = double(4) = 8
double_after_inc = compose(double, increment)
print(double_after_inc(3))  # 8

# increment(double(3)) = increment(6) = 7
inc_after_double = compose(increment, double)
print(inc_after_double(3))  # 7

管道模式(Pipeline)

def pipeline(value, *funcs):
    """将 value 依次通过所有函数处理"""
    for func in funcs:
        value = func(value)
    return value

def strip(s):
    return s.strip()

def capitalize(s):
    return s.capitalize()

def add_period(s):
    return s + "."

result = pipeline("  hello world  ", strip, capitalize, add_period)
print(result)  # "Hello world."

数据转换中的函数组合

from functools import reduce

# 更通用的 compose —— 支持任意多个函数
def compose_all(*funcs):
    """组合多个函数"""
    def composed(x):
        for func in reversed(funcs):
            x = func(x)
        return x
    return composed

def to_upper(s):
    return s.upper()

def add_exclamation(s):
    return s + "!"

def repeat(s):
    return s + " " + s

transform = compose_all(add_exclamation, to_upper, repeat)
print(transform("hello"))  # "HELLO HELLO!"

完整示例:数据清洗管道

下面综合运用本章的所有内容,构建一个数据处理管道:

from functools import partial, reduce

# 1. 数据
raw_data = [
    "  Alice,28,Engineer ",
    "Bob,35,Designer",
    "  Charlie,22,Student",
    "david,29,Data Scientist",
    "",
    "  Eve,31,Manager ",
]

# 2. 处理函数
def clean_row(row):
    """清洗单行数据"""
    row = row.strip()
    if not row:
        return None
    parts = [p.strip() for p in row.split(",")]
    if len(parts) != 3:
        return None
    name, age, job = parts
    return {
        "name": name.capitalize(),
        "age": int(age),
        "job": job,
    }

def is_valid(record):
    """过滤无效记录"""
    return record is not None

def is_adult(record):
    """筛选成年人"""
    return record["age"] >= 18

def format_output(record):
    """格式化输出"""
    return f"{record['name']} ({record['age']}) - {record['job']}"

# 3. 管道
clean_data = list(filter(is_valid, map(clean_row, raw_data)))
adults = filter(is_adult, clean_data)
result = list(map(format_output, adults))

for line in result:
    print(line)
# Alice (28) - Engineer
# Bob (35) - Designer
# Charlie (22) - Student
# David (29) - Data Scientist
# Eve (31) - Manager

实战练习

练习 1:自定义排序

from operator import itemgetter

# 按优先级排序:status(active 优先),然后是名字
tasks = [
    {"name": "修复 bug", "status": "active", "priority": 1},
    {"name": "编写文档", "status": "done", "priority": 3},
    {"name": "添加功能", "status": "active", "priority": 2},
    {"name": "代码审查", "status": "pending", "priority": 1},
]

# 定义排序键函数
def sort_key(task):
    status_order = {"active": 0, "pending": 1, "done": 2}
    return (status_order[task["status"]], task["priority"])

sorted_tasks = sorted(tasks, key=sort_key)
for t in sorted_tasks:
    print(f"[{t['status']:7s}] {t['name']} (优先级 {t['priority']})")

练习 2:从数据中提取唯一值

from functools import reduce

records = [
    {"city": "北京", "temp": 28},
    {"city": "上海", "temp": 32},
    {"city": "北京", "temp": 30},
    {"city": "广州", "temp": 35},
    {"city": "上海", "temp": 31},
]

# 用 reduce 收集不重复的城市
unique_cities = reduce(
    lambda acc, record: acc if record["city"] in acc else acc + [record["city"]],
    records,
    []
)
print(unique_cities)  # ['北京', '上海', '广州']

小结

本章我们学习了 Python 中的函数式编程工具:

  • 高阶函数:接收或返回函数的函数
  • Lambda 表达式:创建简单匿名函数的快捷方式
  • mapfilterreduce:对集合进行批量操作的三剑客
  • sorted 的 key 参数:灵活的自定义排序
  • operator 模块:用 itemgetterattrgetter 等替代简单 lambda
  • functools.partial:冻结函数参数,创建偏函数
  • 函数组合:通过管道和组合构建清晰的数据处理链

下一步: 当程序变复杂时,我们需要把代码组织到多个文件中。下一章我们将学习 Python 的模块和包管理机制。

Summary: 高阶函数、Lambda、map/filter/reduce 与偏函数。