如何定义带前置可选参数的Python函数?以torch.randint为例
torch.randint) Great question! I’ve wondered about this myself when using PyTorch’s APIs. Normally, Python throws an error if you put an optional parameter (with a default value) before a required one—like def func(a=0, b): ... is invalid. But PyTorch gets around this with position-only parameters and keyword-only parameters. Let’s break it down step by step.
The Secret: Python's / and * Parameter Separators
PyTorch leverages two Python 3 features to create this intuitive signature:
/(Position-only parameters): Parameters before/can only be passed by position, not by keyword. This lets us set a default for the first parameter without causing ambiguity with subsequent required parameters.*(Keyword-only parameters): Parameters after*must be passed by keyword, not by position. This keeps optional advanced parameters (likegeneratorordtype) from cluttering up the positional argument list.
Example: Build Your Own randint-Style Function
Here's how you can replicate PyTorch's signature in your own code (requires Python 3.8+):
def my_randint(low=0, /, high, size, *, generator=None, dtype=None, device=None): """Generate random integers between low (inclusive) and high (exclusive).""" # Example logic to show parameter values print(f"low: {low}, high: {high}, size: {size}") print(f"Optional kwargs: generator={generator}, dtype={dtype}, device={device}") # Add your actual implementation here
Valid Ways to Call This Function
- Omit
low(uses default0):my_randint(10, (3, 3)) # Output: low=0, high=10, size=(3,3) - Explicitly set
lowvia position:my_randint(5, 10, (3, 3)) # Output: low=5, high=10, size=(3,3) - Use keywords for parameters between
/and*(optional):my_randint(5, high=10, size=(3,3)) # Also valid - Pass keyword-only parameters (must use keywords):
my_randint(10, (3,3), dtype="int32", device="cuda")
Invalid Calls (As Expected)
- Trying to pass
lowvia keyword (since it's position-only):my_randint(low=5, 10, (3,3)) # SyntaxError! - Trying to pass keyword-only parameters via position:
my_randint(10, (3,3), None, "int32") # TypeError! generator must be passed as a keyword
Compatibility with Older Python Versions
If you need to support Python versions before 3.8, you can manually parse arguments using *args and **kwargs (though it's less clean):
def my_randint(*args, **kwargs): # Parse positional arguments if len(args) == 2: low = 0 high, size = args elif len(args) == 3: low, high, size = args else: raise ValueError("Expected 2 or 3 positional arguments") # Parse keyword-only parameters generator = kwargs.get("generator") dtype = kwargs.get("dtype") device = kwargs.get("device") # Same logic as before print(f"low: {low}, high: {high}, size: {size}") print(f"Optional kwargs: generator={generator}, dtype={dtype}, device={device}")
Why PyTorch Uses This Pattern
This signature balances usability and clarity:
- Users can quickly call the function with just
highandsize(the most common use case) without typinglow=0every time. - Advanced parameters are kept as keywords, so they don't confuse new users but are accessible when needed.
内容的提问来源于stack exchange,提问作者javadr

