use ndarray::{Array, Array1, Array2, Array3, ArrayBase, Dim, Dimension, OwnedRepr}; use serde::Deserialize; use std::collections::HashMap; use std::fmt::Debug; use std::fs; use zerovec::ZeroVec;
let pad = ((k - 1) * dilation) / 2; let weights: Vec<f32> = w.as_slice().iter().collect(); let bias: Vec<f32> = b.as_slice().iter().collect(); letmut acc = vec![0.0f32; cout];
for t in0..l { // acc = np.zeros((Cout,), dtype=x.dtype)
acc.fill(0.0f32); for ki in0..k { let s = { let s = t as isize + (ki * dilation) as isize - pad as isize; if s < 0 || s >= l as isize { continue;
}
s as usize
}; // can optimise if minibatch guarantees no padding needed let x_src = x.submatrix::<1>(s).unwrap().as_slice(); // x_pad[idx] let base = ki * cin * cout; // kernel[k] for (o, acc_o) in acc.iter_mut().enumerate() { // equivalent to np.matmul() letmut sum = 0.0f32; for (c, &xi) in x_src.iter().enumerate() {
sum += xi * weights[base + c * cout + o];
}
*acc_o += sum;
}
} let start = t * cout; let end = start + cout; let dst_row: &mut [f32] = &mut y.as_mut_slice()[start..end]; for (o, dst) in dst_row.iter_mut().enumerate() { let v = acc[o] + bias[o];
*dst = if v > 0.0 { v } else { 0.0 };
}
}
}
fn elementwise_max(
a: MatrixBorrowed<'_, 2>,
b: MatrixBorrowed<'_, 2>, mut y: MatrixBorrowedMut<'_, 2>,
) { let (l, c) = y.as_borrowed().dim(); let aslice = a.as_slice(); let bslice = b.as_slice(); let dst = y.as_mut_slice();
for t in0..l { let start = t * c; let end = start + c;
let ar = &aslice[start..end]; let br = &bslice[start..end]; let dr = &mut dst[start..end];
for o in0..c { let av = ar[o]; let bv = br[o];
dr[o] = if av > bv { av } else { bv };
}
}
}
let w_zeroslice = w.as_slice(); let b_zeroslice = b.as_slice();
letmut acc = vec![0.0f32; classes_usize];
for t in0..l { // logits = x_row * w + b
acc.fill(0.0); let x_row = x.submatrix::<1>(t).unwrap().as_slice();
// matmul for (o, acc_o) in acc.iter_mut().enumerate() { letmut sum = 0.0f32; for (c, &xi) in x_row.iter().enumerate() { let w_val = w_zeroslice.get(c * classes_usize + o).unwrap();
sum += xi * w_val;
}
sum += b_zeroslice.get(o).unwrap();
*acc_o = sum;
}
// logits -= max(logits) letmut m = acc[0]; for &v in &acc[1..] { if v > m {
m = v;
}
} for v in acc.iter_mut() {
*v -= m;
}
// probs = exp(logits); probs /= sum(probs) letmut s = 0.0f32; for v in acc.iter_mut() {
*v = v.exp();
s += *v;
} for v in acc.iter_mut() {
*v *= 1.0 / s;
}
let start = t * classes_usize; let end = start + classes_usize;
y.as_mut_slice()[start..end].copy_from_slice(&acc);
}
}
fn argmax(slice: &[f32]) -> usize { letmut bi = 0usize; letmut bv = slice[0]; for (i, &v) in slice.iter().enumerate().skip(1) { if v > bv {
bv = v;
bi = i;
}
}
bi
}
#[test] fn main() -> Result<(), Box<dyn std::error::Error>> { let path = "tests/cnn/sample.json".to_string(); let json = fs::read_to_string(&path)?;
let rawcnndata: RawCnnData = serde_json::from_str(&json)?; let cnndata = rawcnndata
.try_convert()
.map_err(|_| "validation/conversion failed".to_string())?; let segmenter = CnnSegmenter::new(&cnndata);
let thai = "ปัญหาความแตกต่างที่เกิดขึ้นระหว่างความเป็นธรรมในทางสังคมกับความเป็นธรรมทางกฎหมาย".to_string();
println!("Input: {}", thai); let out = segmenter.segment_str(&thai);
let (l, c) = out.bies.probs.as_borrowed().dim(); let flat = out.bies.probs.as_borrowed().as_slice();
letmut tags = String::with_capacity(l); for t in0..l { let row = &flat[t * c..(t + 1) * c]; let idx = argmax(row);
tags.push(class_idx_to_bies(idx));
}
Die Informationen auf dieser Webseite wurden
nach bestem Wissen sorgfältig zusammengestellt. Es wird jedoch weder Vollständigkeit, noch Richtigkeit,
noch Qualität der bereit gestellten Informationen zugesichert.
Bemerkung:
Die farbliche Syntaxdarstellung und die Messung sind noch experimentell.