mirror of
https://github.com/huggingface/candle.git
synced 2025-06-17 11:08:52 +00:00
Add a hack for generating random uniform/normal for f16/bf16. (#1228)
This commit is contained in:
@ -185,8 +185,14 @@ impl Device {
|
|||||||
Ok(Storage::Cpu(storage))
|
Ok(Storage::Cpu(storage))
|
||||||
}
|
}
|
||||||
Device::Cuda(device) => {
|
Device::Cuda(device) => {
|
||||||
let storage = device.rand_uniform(shape, dtype, lo, up)?;
|
// TODO: Remove the special case if we start supporting generating f16/bf16 directly.
|
||||||
Ok(Storage::Cuda(storage))
|
if dtype == DType::F16 || dtype == DType::BF16 {
|
||||||
|
let storage = device.rand_uniform(shape, DType::F32, lo, up)?;
|
||||||
|
Storage::Cuda(storage).to_dtype(&crate::Layout::contiguous(shape), dtype)
|
||||||
|
} else {
|
||||||
|
let storage = device.rand_uniform(shape, dtype, lo, up)?;
|
||||||
|
Ok(Storage::Cuda(storage))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -213,8 +219,14 @@ impl Device {
|
|||||||
Ok(Storage::Cpu(storage))
|
Ok(Storage::Cpu(storage))
|
||||||
}
|
}
|
||||||
Device::Cuda(device) => {
|
Device::Cuda(device) => {
|
||||||
let storage = device.rand_normal(shape, dtype, mean, std)?;
|
// TODO: Remove the special case if we start supporting generating f16/bf16 directly.
|
||||||
Ok(Storage::Cuda(storage))
|
if dtype == DType::F16 || dtype == DType::BF16 {
|
||||||
|
let storage = device.rand_normal(shape, DType::F32, mean, std)?;
|
||||||
|
Storage::Cuda(storage).to_dtype(&crate::Layout::contiguous(shape), dtype)
|
||||||
|
} else {
|
||||||
|
let storage = device.rand_normal(shape, dtype, mean, std)?;
|
||||||
|
Ok(Storage::Cuda(storage))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
Reference in New Issue
Block a user