mirror of
https://github.com/huggingface/candle.git
synced 2025-06-16 10:38:54 +00:00
Add some group parameter to convolutions. (#566)
* Add some group parameter to convolutions. * Avoid some unnecessary groups checks. * Move the tensor convolution bits. * Properh handling of groups. * Bump the crate version. * And add a changelog.
This commit is contained in:
@ -9,8 +9,8 @@ categories.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
candle = { path = "../../candle-core", version = "0.1.2", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.2" }
|
||||
candle = { path = "../../candle-core", version = "0.1.3", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.3" }
|
||||
num-traits = { workspace = true }
|
||||
tokenizers = { workspace = true, features = ["unstable_wasm"] }
|
||||
|
||||
|
@ -9,8 +9,8 @@ categories.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
candle = { path = "../../candle-core", version = "0.1.2", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.2" }
|
||||
candle = { path = "../../candle-core", version = "0.1.3", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.3" }
|
||||
num-traits = { workspace = true }
|
||||
tokenizers = { workspace = true, features = ["unstable_wasm"] }
|
||||
|
||||
|
@ -295,10 +295,12 @@ impl AudioEncoder {
|
||||
let cfg1 = Conv1dConfig {
|
||||
padding: 1,
|
||||
stride: 1,
|
||||
groups: 1,
|
||||
};
|
||||
let cfg2 = Conv1dConfig {
|
||||
padding: 1,
|
||||
stride: 2,
|
||||
groups: 1,
|
||||
};
|
||||
let conv1 = conv1d(cfg.num_mel_bins, n_state, 3, cfg1, vb.pp("conv1"))?;
|
||||
let conv2 = conv1d(n_state, n_state, 3, cfg2, vb.pp("conv2"))?;
|
||||
|
@ -9,8 +9,8 @@ categories.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
candle = { path = "../../candle-core", version = "0.1.2", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.2" }
|
||||
candle = { path = "../../candle-core", version = "0.1.3", package = "candle-core" }
|
||||
candle-nn = { path = "../../candle-nn", version = "0.1.3" }
|
||||
num-traits = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
|
@ -97,7 +97,11 @@ impl ConvBlock {
|
||||
padding: Option<usize>,
|
||||
) -> Result<Self> {
|
||||
let padding = padding.unwrap_or(k / 2);
|
||||
let cfg = Conv2dConfig { padding, stride };
|
||||
let cfg = Conv2dConfig {
|
||||
padding,
|
||||
stride,
|
||||
groups: 1,
|
||||
};
|
||||
let conv = conv2d_no_bias(c1, c2, k, cfg, vb.pp("conv"))?;
|
||||
let bn = batch_norm(c2, 1e-3, vb.pp("bn"))?;
|
||||
Ok(Self { conv, bn })
|
||||
|
Reference in New Issue
Block a user