use candle_core::{Device, Tensor, Result, Module};
use rlkit::network::NeuralNetwork;
#[test]
fn test_qnetwork_creation() -> Result<()> {
let device = Device::Cpu;
let input_dim = 4;
let hidden_dims = &[64, 32];
let output_dim = 2;
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
assert_eq!(network.hidden_dims().len(), hidden_dims.len());
assert_eq!(network.output_dim(), output_dim);
assert_eq!(network.input_dim(), input_dim);
Ok(())
}
#[test]
fn test_qnetwork_forward_pass() -> Result<()> {
let device = Device::Cpu;
let input_dim = 4;
let hidden_dims = &[8, 4];
let output_dim = 2;
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
let input = Tensor::randn(0f32, 1f32, (2, input_dim), &device)?;
let output = network.forward(&input)?;
let output_shape = output.shape().dims();
assert_eq!(output_shape.len(), 2);
assert_eq!(output_shape[0], 2); assert_eq!(output_shape[1], output_dim);
Ok(())
}
#[test]
fn test_qnetwork_parameters() -> Result<()> {
let device = Device::Cpu;
let input_dim = 4;
let hidden_dims = &[8, 4];
let output_dim = 2;
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
let params = network.parameters();
assert_eq!(params.len(), 6);
let shape_0 = params[0].shape().dims();
assert_eq!(shape_0[0], 8);
assert_eq!(shape_0[1], 4);
let shape_1 = params[1].shape().dims();
assert_eq!(shape_1[0], 8);
let shape_2 = params[2].shape().dims();
assert_eq!(shape_2[0], 4);
assert_eq!(shape_2[1], 8);
let shape_3 = params[3].shape().dims();
assert_eq!(shape_3[0], 4);
let shape_4 = params[4].shape().dims();
assert_eq!(shape_4[0], 2);
assert_eq!(shape_4[1], 4);
let shape_5 = params[5].shape().dims();
assert_eq!(shape_5[0], 2);
Ok(())
}
#[test]
fn test_qnetwork_save_and_load() -> Result<()> {
use std::fs;
let device = Device::Cpu;
let input_dim = 4;
let hidden_dims = &[8, 4];
let output_dim = 2;
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
let temp_path = "temp_qnetwork.safetensors";
network.save(temp_path)?;
let loaded_network = NeuralNetwork::load(temp_path, input_dim, hidden_dims, output_dim, &device)?;
let original_params = network.parameters();
let loaded_params = loaded_network.parameters();
assert_eq!(original_params.len(), loaded_params.len());
for (orig, loaded) in original_params.iter().zip(loaded_params.iter()) {
assert_eq!(orig.shape().dims(), loaded.shape().dims());
}
let test_input = Tensor::randn(0f32, 1f32, (1, input_dim), &device)?;
let original_output = network.forward(&test_input)?;
let loaded_output = loaded_network.forward(&test_input)?;
let output_shape = loaded_output.shape().dims();
assert_eq!(output_shape.len(), 2);
assert_eq!(output_shape[0], 1);
assert_eq!(output_shape[1], output_dim);
let diff = original_output.sub(&loaded_output)?.abs()?.mean_all()?;
println!("原始输出: {:?}", original_output.to_vec2::<f32>()?);
println!("加载输出: {:?}", loaded_output.to_vec2::<f32>()?);
let diff_value = diff.to_scalar::<f32>()?;
assert!(diff_value < 1e-5, "加载的模型输出与原始模型输出差异过大: {}", diff_value);
fs::remove_file(temp_path).ok();
Ok(())
}
#[test]
fn test_qnetwork_with_different_dimensions() -> Result<()> {
let device = Device::Cpu;
let test_configs = [
(10, &[128, 64, 32][..], 5), (2, &[16][..], 1), (100, &[256, 128, 64, 32][..], 10), ];
for (input_dim, hidden_dims, output_dim) in test_configs {
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
let input = Tensor::randn(0f32, 1f32, (3, input_dim), &device)?;
let output = network.forward(&input)?;
let output_shape = output.shape().dims();
assert_eq!(output_shape.len(), 2);
assert_eq!(output_shape[0], 3);
assert_eq!(output_shape[1], output_dim);
}
Ok(())
}
#[test]
fn test_qnetwork_training_fitting() -> Result<()> {
use candle_nn::{optim::AdamW, Optimizer};
let device = Device::Cpu;
let input_dim = 1;
let hidden_dims = &[64, 32];
let output_dim = 1;
let network = NeuralNetwork::new(input_dim, hidden_dims, output_dim, &device)?;
let mut optimizer = AdamW::new_lr(network.varmap.all_vars(), 1e-3)?;
let x_train = Tensor::randn(0f32, 1f32, (100, input_dim), &device)?;
let y_train = x_train.powf(2.0)?;
let epochs = 100;
let mut best_loss = f32::INFINITY;
for epoch in 0..epochs {
let y_pred = network.forward(&x_train)?;
let loss = y_pred.sub(&y_train)?.sqr()?.mean_all()?;
let loss_value = loss.to_scalar::<f32>()?;
optimizer.backward_step(&loss)?;
if loss_value < best_loss {
best_loss = loss_value;
}
if epoch % 10 == 0 {
println!("Epoch {}/{}, Loss: {:.6}", epoch + 1, epochs, loss_value);
}
}
println!("最终训练损失: {:.6}", best_loss);
let test_inputs = Tensor::new(&[[-1.0f32], [0.0], [1.0], [2.0]], &device)?;
let expected_outputs = Tensor::new(&[[1.0f32], [0.0], [1.0], [4.0]], &device)?;
let actual_outputs = network.forward(&test_inputs)?;
println!("测试输入: {:?}", test_inputs.to_vec2::<f32>()?);
println!("期望输出: {:?}", expected_outputs.to_vec2::<f32>()?);
println!("实际输出: {:?}", actual_outputs.to_vec2::<f32>()?);
let test_loss = actual_outputs.sub(&expected_outputs)?.sqr()?.mean_all()?;
let test_loss_value = test_loss.to_scalar::<f32>()?;
println!("测试损失: {:.6}", test_loss_value);
assert!(test_loss_value < 0.2, "网络拟合性能不足,测试损失太大: {}", test_loss_value);
Ok(())
}