[Feature request] `ArrayRef<A, Ix2>.dot()` for axis greater than `Ix2`
評估
研究方向
Start from the ndarray dot() entry point and inspect how dimensions are currently constrained. Compare the requested behavior with the provided Ix3-by-Ix2 example, then verify that higher-rank left-hand arrays produce the expected NumPy-compatible shape while preserving existing Ix2 behavior.
由索引模型根據 Issue 內容生成。
描述
In NumPy, the left-hand side of a matrix multiplication can have as many axes as desired, as long as it has more than 2 axes and the last axis's dimension matches that of the 0th axis of the right-hand side, e.g.:
import numpy as np
x = np.random.random((3, 2, 5, 9, 12))
y = np.random.random((12, 13))
(x @ y).shape
# (3, 2, 5, 9, 13)
In ndarray, you can't do this directly:
// Doesn't compile:
use ndarray::prelude::*;
use ndarray_rand::RandomExt;
use ndarray_rand::rand_distr::Uniform;
fn main() {
let x: Array<f64, Ix3> = Array::random(
(12, 4, 3),
Uniform::new(0., 1.).unwrap()
);
let y: Array<f64, Ix2> = Array::random(
(3, 2),
Uniform::new(0., 1.).unwrap()
);
let x_y = x.dot(&y);
println!("{}", x_y);
}
Compiler Output
$ cargo run
Compiling playground v0.1.0 (/home/connor/RustroverProjects/playground)
error[E0275]: overflow evaluating the requirement `&ArrayBase<_, _, _>: Not`
--> src/main.rs:15:17
|
15 | let x_y = x.dot(&y);
| ^^^
|
= help: consider increasing the recursion limit by adding a `#![recursion_limit = "256"]` attribute to your crate (`playground`)
= note: required for `&ArrayBase<_, _, _>` to implement `Not`
= note: 127 redundant requirements hidden
= note: required for `&ArrayBase<OwnedRepr<f64>, Dim<[usize; 3]>, f64>` to implement `Not`
For more information about this error, try `rustc --explain E0275`.
error: could not compile `playground` (bin "playground") due to 1 previous error
Emulating the behavior in the previous NumPy example requires a non-trivial amount of work, e.g. for x with 3 axes:
use ndarray::prelude::*;
use ndarray_rand::RandomExt;
use ndarray_rand::rand_distr::Uniform;
fn main() {
let x: Array<f64, Ix3> = Array::random((12, 4, 3), Uniform::new(0., 1.).unwrap());
let y = Array::random((3, 2), Uniform::new(0., 1.).unwrap());
let (a, b, c) = (x.len_of(Axis(0)), x.len_of(Axis(1)), x.len_of(Axis(2)));
let d = y.len_of(Axis(1));
let x_y: Array3<f64> = x
.to_shape((a * b, c)).unwrap()
.dot(&y)
.to_shape((a, b, d)).unwrap()
.to_owned();
println!("{:?}", x_y);
}
Therefore, I think it would be nice to have dot() be implemented for axis numbers greater than Ix2
- 主要語言
- Rust
- 星號
- 4.3k
- 分支
- 391
- PR 合併指標
- 30 天內沒有已合併 PR
環境準備
這個專案沒有提供開發容器、Dockerfile 或貢獻指南,環境需要你自己搭建:先看它的 README,通用步驟見我們的新手貢獻指南。
從這裡開始
- 先讀完整個 Issue,再讀專案的貢獻指南。
- 在 Issue 下留言說明你要接手 —— 這能避免兩個人做同樣的事。
- Fork 儲存庫,在一個分支上完成修改。
- 送出 Pull Request,並在描述裡引用這個 Issue 編號。
rust-ndarray/ndarray 的其他 Issue
-
.is_all_nan() returns true for an all-finite array可能已有人在做 @youdie006 於 56 天前認領。 未關閉
難度 2/5 1-3 小時 新手友好度 72/100
rust-ndarray/ndarray#1612 · 1 則留言 · 1 個 reaction ·
-
難度 4/5 3-5 天 新手友好度 48/100
rust-ndarray/ndarray#1617 · 1 則留言 ·
-
Stack overflow in `triu`可能已有人在做 @cestercian 於 30 天前認領。 未關閉bug good first issue
難度 3/5 1-2 天 新手友好度 68/100
rust-ndarray/ndarray#1615 · 1 則留言 ·
-
難度 4/5 3-5 天 新手友好度 48/100
rust-ndarray/ndarray#1610 ·
-
Empty Array lead to "The strides must not allow any element to be referenced by two different indices"可能已有人在做 @Guflly 於 68 天前認領。 未關閉
難度 3/5 1-2 天 新手友好度 72/100
rust-ndarray/ndarray#1609 ·
查看 rust-ndarray/ndarray 的全部 Issue
相似的 Issue
-
[Feature]: [P3] engine-rs: the package source hash should ignore line endings and untracked files未關閉
難度 2/5 1-3 小時 新手友好度 70/100
maniator/verticopolis#880 ·
維護者通常 1 天內回覆
-
IO.get_env on Node truncates names at embedded NUL可能已有人在做 @Yi-111-a 今天認領。 未關閉
難度 2/5 1-3 小時 新手友好度 82/100
HigherOrderCO/Bend#1449 · 1 則留言 ·
-
難度 2/5 1-3 小時 新手友好度 82/100
維護者通常 1 天內回覆
-
documentation
難度 2/5 1-3 小時 新手友好度 66/100
維護者通常 3 天內回覆
-
難度 2/5 1-3 小時 新手友好度 62/100
維護者通常 1 天內回覆