跳到正文

Python 函数定义:参数、返回值与作用域

分析脚本写长了必然遇到重复代码:同一段数据清洗在三个 notebook 里各有一份,改一处忘两处。函数是 Python 里最基本的复用单位,把「一次做对的事」固定下来,后面只剩调用。写函数的门槛不高,但参数传递和作用域这两块不清楚的话,脚本会从「能跑」变成「不知道为什么会跑出这个结果」。

def add(a, b):
return a + b
print(add(2, 3))
print(add(b=3, a=2))
5
5

def 后面跟函数名和参数列表,冒号换行后是函数体。调用时参数可以按位置传,也可以按名字传——按名字传更啰嗦,但参数一多就不用去数位置了。

函数体里没有 return,或者 return 后面什么都不写,返回值是 None

def log_it(x):
print(f"处理 {x}")
result = log_it(5)
print(result)
处理 5
None

log_it(5) 照样打印了那行提示,但 result 拿到的是 None。忘了写 return 是新手最常见的错误,表现是下游拿到 None 之后才报错,位置离真正的源头很远。

需要返回多个值时,直接 return 一个元组,调用方解包:

def minmax(x):
return min(x), max(x)
lo, hi = minmax([3, 1, 4, 1, 5])
print(lo, hi)
t = minmax([3, 1, 4, 1, 5])
print(t)
1 5
(1, 5)

Python 没有 R 那种「最后一行表达式的值就是返回值」的规则,return 必须显式写出来。函数遇到 return 立即结束,后面的代码不会执行——提前退出的写法与别的语言一致:

import math
def safe_log(x):
if x <= 0:
return None
return math.log(x)

默认值在函数定义时求值一次,之后每次调用都复用同一个对象。参数是可变对象时,这会直接导致数据串台:

def append_to(x, target=[]):
target.append(x)
return target
print(append_to(1))
print(append_to(2))
print(append_to(3, []))
[1]
[1, 2]
[3]

第二次调用返回 [1, 2],第一次追加进去的 1 还在。因为默认的 [] 只在 def 执行时创建了一次,两次调用共用同一个列表。第三次显式传了 [],才拿到干净的结果。

运行时求值的直觉在这里会误导人。下面这个例子把这件事摊开:

import time
def stamp(t=time.time()):
return t
a = stamp()
time.sleep(0.05)
b = stamp()
print(a == b)
True

两次调用隔了 0.05 秒,返回值完全相同。time.time() 在定义函数时就调用了,不是每次调用时重算。

修法是默认值用 None,在函数体内部再创建:

def append_to(x, target=None):
if target is None:
target = []
target.append(x)
return target
print(append_to(1))
print(append_to(2))
[1]
[2]

不可变对象(数字、字符串、元组、None)没有这个问题,所以 def f(x, n=10) 这种写法是安全的。判断标准是默认值本身会不会被修改。

*args 把所有多出来的位置参数收成一个元组,**kwargs 把多出来的关键字参数收成一个字典:

def summarize(*args):
print(f"收到 {len(args)} 个位置参数: {args}")
summarize(1, 2, 3)
summarize("a", "b")
收到 3 个位置参数: (1, 2, 3)
收到 2 个位置参数: ('a', 'b')
def configure(**kwargs):
for k, v in kwargs.items():
print(f"{k} = {v}")
configure(alpha=0.05, method="tukey")
alpha = 0.05
method = tukey

单个 * 后面跟的命名参数只能按关键字传,位置参数传不进去:

import numpy as np
def normalize(x, *, center=True, scale=True):
arr = np.asarray(x, dtype=float)
if center:
arr = arr - arr.mean()
if scale:
arr = arr / arr.std(ddof=1)
return arr
print(np.round(normalize([1, 2, 3, 4]), 3))
[-1.162 -0.387 0.387 1.162]

第二个参数写成位置传会直接报错:

normalize([1, 2, 3, 4], True)
TypeError: normalize() takes 1 positional argument but 2 were given

这种「仅关键字参数(keyword-only argument)」在参数一多时能防住顺序传错。np.roundsortedpd.read_csv 里大量使用这个设计,pd.read_csv(path, True) 这种写法会立刻失败,而不是静默地把 True 当成别的选项。

Python 里函数是一等对象,可以赋给变量、放进容器、作为参数传给别的函数:

def square(x):
return x * x
def cube(x):
return x ** 3
funcs = {"square": square, "cube": cube}
print(funcs["cube"](3))
print(sorted([3, 1, 2], key=lambda v: -v))
27
[3, 2, 1]

sortedkey 参数收的就是一个函数对象,lambda v: -v 是匿名函数的写法。用字典把名字映射到函数,可以替代一长串 if/elif 分支。

Python 的传参规则常被概括成「传对象引用」。实际后果只有两条,记住这两条就够了。

函数内部修改可变对象(列表、字典、DataFrame),调用方看得见:

def modify(v):
v[0] = 999
return v
z = [1, 2, 3]
print(modify(z))
print(z)
[999, 2, 3]
[999, 2, 3]

z 被改了。这一点与 R 相反:R 的函数调用是传值语义,函数内改向量不会影响外面。从 R 转过来的人在这里栽过不少次。

函数内部重新绑定名字,调用方看不见:

def rebind(v):
v = [999]
return v
z2 = [1, 2, 3]
print(rebind(z2))
print(z2)
[999]
[1, 2, 3]

v = [999] 只是让局部名字 v 指向了一个新列表,原来的 z2 没动。区分「改内容」和「改指向」是理解 Python 参数传递的关键。

写函数时如果不打算修改传入的对象,就不要在函数体里对它做原地操作。df.drop(columns=[...]) 默认返回新对象,df.drop(columns=[...], inplace=True) 会改原表——后者在新版本 pandas 里已经不推荐使用。

Python 查找一个变量名,按四层顺序找:局部(Local)→ 外层函数(Enclosing)→ 全局(Global)→ 内置(Built-in),缩写 LEGB。

x = "global"
def outer():
x = "enclosing"
def inner():
x = "local"
return x
return inner()
print(outer())
print(x)
local
global

inner 在自己这层找到了 x 就返回,不再往外找。三层的名字互不干扰,函数结束后局部的 x 就没了。

查找路径由函数定义的位置决定,不由调用位置决定:

x = 100
def f():
return x
def g():
x = 1
return f()
print(g())
100

g() 里明明有 x = 1f() 还是返回 100。因为 f 定义在全局作用域,它去全局找 x。这个规则叫词法作用域(lexical scoping),与 R 一致。

函数内部给一个外层的名字赋值,Python 会认为你要创建局部变量。想改全局变量得用 global 声明:

total = 0
def acc(vals):
global total
total += sum(vals)
return total
print(acc([1, 2, 3]))
print(acc([1, 2, 3]))
print(total)
6
12
12

每次调用都把全局的 total 改掉。忘了写 global 不会静默出错,而是直接抛异常:

total2 = 0
def acc2(vals):
total2 = total2 + sum(vals)
return total2
acc2([1, 2, 3])
UnboundLocalError: cannot access local variable 'total2' where it is not associated with a value

报错的原因是:只要函数体里出现了对 total2 的赋值,Python 就把整个函数里的 total2 都当成局部变量;右侧读取时该局部变量还没绑定,于是报错。R 里对应的写法只会给出错误结果而不报错,Python 这里直接失败,排查起来更快。

修改外层函数(非全局)的变量用 nonlocal

def counter():
n = 0
def step():
nonlocal n
n += 1
return n
return step
c = counter()
print(c(), c(), c())
1 2 3

counter() 返回的 step 记住了 n,这种结构叫闭包(closure)。global 会一直影响模块级的变量,能用闭包或返回值解决的问题,优先不要用 global:同样调用两次 acc([1, 2, 3]),第一次返回 6、第二次返回 12,结果取决于调用历史,测试时没法单独验证某一次调用。

闭包捕获的是变量本身,不是当时的值。在循环里建函数,容易踩到这个坑:

fs = [lambda: i for i in range(3)]
print([fn() for fn in fs])
[2, 2, 2]

三个函数都返回 2,因为循环结束时 i 就是 2,而它们都引用同一个 i。把 i 绑定成默认参数可以固定住当时的值:

fs2 = [lambda i=i: i for i in range(3)]
print([fn() for fn in fs2])
[0, 1, 2]

默认参数在定义时求值,正好把当前的值存了下来。这个写法看着别扭,但它是标准解法。

lambda 只能写一个表达式,不能含赋值、for 循环、try 语句。需要这些就得写 def。给 sortedmaxmap 传一个短小的键函数时用 lambda 合适,超过一行就换成命名函数——栈里显示 lambda 而不是函数名,调试时不好定位。

类型注解(type annotation)写在参数和返回值上,Python 运行时不会强制检查,但编辑器能据此提示,阅读代码的人也能一眼看出接口约定:

import numpy as np
def cv(values: list[float]) -> float:
"""返回变异系数:标准差除以均值。"""
arr = np.asarray(values, dtype=float)
return float(arr.std(ddof=1) / arr.mean())
print(cv([12.4, 11.8, 13.1, 18.6, 17.9]))
print(cv.__doc__)
0.21873083526627332
返回变异系数:标准差除以均值。

三引号包起来的第一段字符串就是 docstring,help(cv)cv.__doc__ 都能读到。项目里约定好在 docstring 里写清参数含义、返回值和特殊约定(比如「缺失值会被忽略」),半年后回来的第一个人就是你。

把上面几件事合起来:参数带默认值、按名字传可选参数、返回字典。该函数从数据里算出样本量、均值、标准差、标准误和 95% 置信区间。

标准误(standard error)是均值这个估计量的标准差,等于样本标准差除以样本量的平方根;95% 置信区间(confidence interval)用 t 分布的分位数算,自由度是 n-1。

import numpy as np
from scipy import stats
def describe(x, na_rm=True, digits=2, conf=0.95):
"""返回样本量、均值、标准差、标准误与 95% 置信区间。"""
values = np.asarray(x, dtype=float)
if na_rm:
values = values[~np.isnan(values)]
n = values.size
mean = values.mean()
sd = values.std(ddof=1)
se = sd / np.sqrt(n)
q = stats.t.ppf(1 - (1 - conf) / 2, df=n - 1)
est = {"mean": mean, "sd": sd, "se": se, "lower": mean - q * se, "upper": mean + q * se}
return {"n": int(n), **{k: round(float(v), digits) for k, v in est.items()}}
ozone = np.array([41, 36, 12, 18, np.nan, 28, 23, 19, 8, 7, 16, 11, 14, 18, 20])
print(describe(ozone))
print(describe(ozone, na_rm=False))
{'n': 14, 'mean': 19.36, 'sd': 9.94, 'se': 2.66, 'lower': 13.62, 'upper': 25.09}
{'n': 15, 'mean': nan, 'sd': nan, 'se': nan, 'lower': nan, 'upper': nan}

values.std(ddof=1) 里的 ddof=1 表示除以 n-1 而不是 n,得到的是样本标准差。NumPy 默认 ddof=0,用于描述性统计时算的是总体标准差,写论文报告标准差时要显式指定。

na_rm=False 那次调用把 NaN 留在了数组里,所有统计量都被污染成 nan。NumPy 的 mean() 与 R 的 mean(na.rm = FALSE) 行为一致,但 NumPy 不提供 na.rm 参数,缺省处理方式是用 np.nanmean() 这类专用函数,或者像本例这样在入口处过滤。

换一个数据集交叉验证:

from sklearn.datasets import load_iris
iris = load_iris()
sepal = iris.data[:, 0]
print(describe(sepal))
print(iris.feature_names)
{'n': 150, 'mean': 5.84, 'sd': 0.83, 'se': 0.07, 'lower': 5.71, 'upper': 5.98}
['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)']

150 个鸢尾花样本的花萼长度均值 5.84 厘米,95% 置信区间 5.71 到 5.98。函数里每一处都只依赖参数:数据通过 x 进来,na_rm 控制缺失值处理,digits 控制输出精度,没有任何全局变量。这样的函数可以直接搬进另一个脚本,也能单独写测试。

R 语言的同名函数用 sd()qt() 组合出同一套逻辑,写法对照见 /r/basics/functions/。函数写完,Python 基础部分就齐了;接下来是 pandas 的数据操作,从 /python/pandas/data-loading/ 开始。

UnboundLocalError: cannot access local variable 'x' 函数体里给 x 赋值了,但右侧又要读它。要么加上 global / nonlocal 声明,要么改用别的变量名。

默认参数在不同调用之间「记住了」上次的值 默认值是可变对象(列表、字典、集合)。改成 None 哨兵,在函数体内创建。

传进函数的列表被意外改了 函数内部做了 v[0] = ...v.sort()v.append(...) 这类原地操作。不需要修改调用方数据时先 v = list(v) 复制一份。

循环里创建的函数结果都一样 闭包捕获的是变量而不是值。用默认参数固定,或者把函数定义抽到一个辅助函数里。

TypeError: f() takes 1 positional argument but 2 were given 函数签名里有 ** 之后的参数只能按关键字传。检查调用处是不是漏写了参数名。