Swift如何实现多维数组沿轴0的最大值索引获取?
argmax(axis=0) in Swift Great question! Let's break down how to get the exact behavior of NumPy's argmax(axis=0) in Swift, since there's no built-in function for this out of the box.
First, let's make sure we're aligned: when using axis=0, we want the row index of the maximum value in each column of our 2D array. For your example [[1,2,3],[4,3,1]], that means checking each column individually to find where the largest number lives:
- Column 0: values are 1 and 4 → max lives at row index 1
- Column 1: values are 2 and 3 → max lives at row index 1
- Column 2: values are 3 and 1 → max lives at row index 0
Which gives the result [1,1,0]—exactly what you get from NumPy. Let's build this in Swift.
Step-by-Step Implementation
Here's a generic, easy-to-follow function that works with any comparable type (Int, Double, etc.):
func argmax<T: Comparable>(_ array: [[T]], axis: Int) -> [Int] { // We'll focus on axis=0 first, since that's what you asked for guard axis == 0 else { fatalError("This implementation only handles axis=0 right now—feel free to extend it for axis=1 if needed!") } // Ensure our 2D array is "rectangular" (all rows have the same length) guard let columnCount = array.first?.count else { return [] } precondition(array.allSatisfy { $0.count == columnCount }, "All rows in the array must have the same number of elements") var resultIndices = [Int]() // Loop through each column for column in 0..<columnCount { // Start with the first row's value as our initial maximum var maxValue = array[0][column] var maxIndex = 0 // Check every other row in this column for row in 1..<array.count { let currentValue = array[row][column] if currentValue > maxValue { maxValue = currentValue maxIndex = row } } // Add the index of the maximum value for this column to our result resultIndices.append(maxIndex) } return resultIndices }
Testing It With Your Example
Let's plug in your sample array to verify it works as expected:
let a = [[1, 2, 3], [4, 3, 1]] let maxIndices = argmax(a, axis: 0) print(maxIndices) // Prints: [1, 1, 0]
Perfect—matches the NumPy output exactly!
Bonus: Works With Other Data Types
Since we used a generic T: Comparable, this function also works seamlessly with floating-point numbers:
let floatArray = [[1.5, 2.2, 3.7], [4.1, 3.9, 1.0]] let floatMaxIndices = argmax(floatArray, axis: 0) print(floatMaxIndices) // Prints: [1, 1, 0]
Performance Note
For small to medium-sized arrays, this native Swift implementation will be totally sufficient. If you're working with huge datasets, you could explore optimizing with the Accelerate framework (Apple's high-performance numerical computing library), but that adds complexity. The above function is a great starting point for most use cases.
内容的提问来源于stack exchange,提问作者gi097

