Force math attention in F32 reference captures
This commit is contained in:
parent
e73e557ee1
commit
0494e631e9
@ -76,6 +76,11 @@ if(EXISTS "${KIMODO_GGML_SOURCE_DIR}/CMakeLists.txt")
|
|||||||
target_include_directories(kimodo-multiprompt-fixture-parity PRIVATE src)
|
target_include_directories(kimodo-multiprompt-fixture-parity PRIVATE src)
|
||||||
target_link_libraries(kimodo-multiprompt-fixture-parity PRIVATE ggml ggml-vulkan)
|
target_link_libraries(kimodo-multiprompt-fixture-parity PRIVATE ggml ggml-vulkan)
|
||||||
target_compile_definitions(kimodo-multiprompt-fixture-parity PRIVATE KIMODO_HAVE_GGML=1 KIMODO_HAVE_GGML_VULKAN=1)
|
target_compile_definitions(kimodo-multiprompt-fixture-parity PRIVATE KIMODO_HAVE_GGML=1 KIMODO_HAVE_GGML_VULKAN=1)
|
||||||
|
add_executable(kimodo-stage-fixture-parity tests/stage_fixture_parity.cpp)
|
||||||
|
target_sources(kimodo-stage-fixture-parity PRIVATE src/gguf.cpp src/ggml_weights.cpp src/denoiser.cpp src/motion_rep.cpp src/diffusion.cpp)
|
||||||
|
target_include_directories(kimodo-stage-fixture-parity PRIVATE src)
|
||||||
|
target_link_libraries(kimodo-stage-fixture-parity PRIVATE ggml ggml-vulkan)
|
||||||
|
target_compile_definitions(kimodo-stage-fixture-parity PRIVATE KIMODO_HAVE_GGML=1 KIMODO_HAVE_GGML_VULKAN=1)
|
||||||
add_executable(kimodo-llm-text-session-parity
|
add_executable(kimodo-llm-text-session-parity
|
||||||
tests/llm_text_session_parity.cpp src/llm_text_encoder.cpp src/llm_tokenizer.cpp)
|
tests/llm_text_session_parity.cpp src/llm_text_encoder.cpp src/llm_tokenizer.cpp)
|
||||||
target_include_directories(kimodo-llm-text-session-parity PRIVATE src)
|
target_include_directories(kimodo-llm-text-session-parity PRIVATE src)
|
||||||
|
|||||||
@ -85,12 +85,29 @@ def main() -> None:
|
|||||||
sys.path.insert(0, str(upstream))
|
sys.path.insert(0, str(upstream))
|
||||||
from kimodo import load_model # pylint: disable=import-outside-toplevel
|
from kimodo import load_model # pylint: disable=import-outside-toplevel
|
||||||
|
|
||||||
|
# The NVIDIA PyTorch container enables TF32 globally. Kimodo GGML uses
|
||||||
|
# F32 accumulation for reference parity, so a CUDA capture must disable
|
||||||
|
# Tensor Core TF32 before any model module is materialised.
|
||||||
|
if args.device.startswith("cuda"):
|
||||||
|
torch.backends.cuda.matmul.allow_tf32 = False
|
||||||
|
torch.backends.cudnn.allow_tf32 = False
|
||||||
|
torch.set_float32_matmul_precision("highest")
|
||||||
|
# TransformerEncoder otherwise selects PyTorch's CUDA fast path,
|
||||||
|
# whose fused attention reductions have a different F32 accumulation
|
||||||
|
# order from Kimodo's explicit GGML attention graph.
|
||||||
|
torch.backends.mha.set_fastpath_enabled(False)
|
||||||
|
torch.backends.cuda.enable_flash_sdp(False)
|
||||||
|
torch.backends.cuda.enable_mem_efficient_sdp(False)
|
||||||
|
torch.backends.cuda.enable_math_sdp(True)
|
||||||
|
|
||||||
torch.manual_seed(args.seed)
|
torch.manual_seed(args.seed)
|
||||||
encoder = CapturedEmbeddings(args.prompt, args.embedding)
|
encoder = CapturedEmbeddings(args.prompt, args.embedding)
|
||||||
model, resolved = load_model("kimodo-smplx-rp", device=args.device,
|
model, resolved = load_model("kimodo-smplx-rp", device=args.device,
|
||||||
text_encoder=encoder, return_resolved_name=True)
|
text_encoder=encoder, return_resolved_name=True)
|
||||||
calls: list[dict[str, list[torch.Tensor] | torch.Tensor]] = []
|
calls: list[dict[str, list[torch.Tensor] | torch.Tensor]] = []
|
||||||
inverse_inputs: list[torch.Tensor] = []
|
inverse_inputs: list[torch.Tensor] = []
|
||||||
|
root_calls: list[tuple[tuple[torch.Tensor, ...], torch.Tensor]] = []
|
||||||
|
body_calls: list[tuple[tuple[torch.Tensor, ...], torch.Tensor]] = []
|
||||||
original_step = model.denoising_step
|
original_step = model.denoising_step
|
||||||
|
|
||||||
def capture_step(*values: Any, **kwargs: Any) -> torch.Tensor:
|
def capture_step(*values: Any, **kwargs: Any) -> torch.Tensor:
|
||||||
@ -110,6 +127,18 @@ def main() -> None:
|
|||||||
call["output"].append(result.detach().clone()) # type: ignore[index]
|
call["output"].append(result.detach().clone()) # type: ignore[index]
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
def capture_stage(calls_for_stage: list[tuple[tuple[torch.Tensor, ...], torch.Tensor]]):
|
||||||
|
def hook(_module: torch.nn.Module, values: tuple[Any, ...], output: torch.Tensor) -> None:
|
||||||
|
if not calls_for_stage:
|
||||||
|
calls_for_stage.append((
|
||||||
|
tuple(value.detach().clone() for value in values if isinstance(value, torch.Tensor)),
|
||||||
|
output.detach().clone(),
|
||||||
|
))
|
||||||
|
return hook
|
||||||
|
|
||||||
|
root_hook = model.denoiser.model.root_model.register_forward_hook(capture_stage(root_calls))
|
||||||
|
body_hook = model.denoiser.model.body_model.register_forward_hook(capture_stage(body_calls))
|
||||||
|
|
||||||
original_inverse = model.motion_rep.inverse
|
original_inverse = model.motion_rep.inverse
|
||||||
|
|
||||||
def capture_inverse(motion: torch.Tensor, *values: Any, **kwargs: Any):
|
def capture_inverse(motion: torch.Tensor, *values: Any, **kwargs: Any):
|
||||||
@ -129,10 +158,17 @@ def main() -> None:
|
|||||||
finally:
|
finally:
|
||||||
model.denoising_step = original_step
|
model.denoising_step = original_step
|
||||||
model.motion_rep.inverse = original_inverse
|
model.motion_rep.inverse = original_inverse
|
||||||
if len(calls) != len(args.prompt) or not inverse_inputs:
|
root_hook.remove()
|
||||||
|
body_hook.remove()
|
||||||
|
if len(calls) != len(args.prompt) or not inverse_inputs or not root_calls or not body_calls:
|
||||||
raise RuntimeError(f"expected {len(args.prompt)} segment trajectories, got {len(calls)}")
|
raise RuntimeError(f"expected {len(args.prompt)} segment trajectories, got {len(calls)}")
|
||||||
|
|
||||||
arrays: dict[str, np.ndarray] = {"stitched_motion_rep": as_f32(inverse_inputs[-1])}
|
arrays: dict[str, np.ndarray] = {"stitched_motion_rep": as_f32(inverse_inputs[-1])}
|
||||||
|
for stage, stage_calls in (("root", root_calls), ("body", body_calls)):
|
||||||
|
values, stage_output = stage_calls[0]
|
||||||
|
arrays[f"{stage}_output"] = as_f32(stage_output)
|
||||||
|
for index, value in enumerate(values):
|
||||||
|
arrays[f"{stage}_input_{index}"] = as_f32(value)
|
||||||
for index, call in enumerate(calls):
|
for index, call in enumerate(calls):
|
||||||
prefix = f"segment_{index:02d}_"
|
prefix = f"segment_{index:02d}_"
|
||||||
for name in ("pad_mask", "text_features", "text_pad_mask", "first_heading_angle", "motion_mask", "observed_motion"):
|
for name in ("pad_mask", "text_features", "text_pad_mask", "first_heading_angle", "motion_mask", "observed_motion"):
|
||||||
@ -158,6 +194,8 @@ def main() -> None:
|
|||||||
"transition_frames": args.transition_frames, "diffusion_steps": args.steps,
|
"transition_frames": args.transition_frames, "diffusion_steps": args.steps,
|
||||||
"seed": args.seed, "cfg_type": "separated", "cfg_weight": [2.0, 2.0],
|
"seed": args.seed, "cfg_type": "separated", "cfg_weight": [2.0, 2.0],
|
||||||
"post_processing": False, "device": args.device, "torch": torch.__version__,
|
"post_processing": False, "device": args.device, "torch": torch.__version__,
|
||||||
|
"cuda_tf32": torch.backends.cuda.matmul.allow_tf32 if args.device.startswith("cuda") else None,
|
||||||
|
"cuda_mha_fastpath": torch.backends.mha.get_fastpath_enabled() if args.device.startswith("cuda") else None,
|
||||||
"python": platform.python_version(), "embedding_sources": [
|
"python": platform.python_version(), "embedding_sources": [
|
||||||
{"path": str(path), "sha256": sha256(path), "note": note}
|
{"path": str(path), "sha256": sha256(path), "note": note}
|
||||||
for path, note in zip(args.embedding, notes)
|
for path, note in zip(args.embedding, notes)
|
||||||
|
|||||||
36
tests/stage_fixture_parity.cpp
Normal file
36
tests/stage_fixture_parity.cpp
Normal file
@ -0,0 +1,36 @@
|
|||||||
|
#include "denoiser.hpp"
|
||||||
|
#include "ggml_weights.hpp"
|
||||||
|
#include <algorithm>
|
||||||
|
#include <cmath>
|
||||||
|
#include <cstdio>
|
||||||
|
#include <fstream>
|
||||||
|
#include <stdexcept>
|
||||||
|
#include <string>
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
static std::vector<float> read(const std::string &path) {
|
||||||
|
std::ifstream input(path, std::ios::binary | std::ios::ate);
|
||||||
|
if (!input || input.tellg() < 0 || input.tellg() % 4) throw std::runtime_error("bad fixture: " + path);
|
||||||
|
std::vector<float> value(static_cast<size_t>(input.tellg()) / 4); input.seekg(0);
|
||||||
|
input.read(reinterpret_cast<char *>(value.data()), static_cast<std::streamsize>(value.size() * 4));
|
||||||
|
if (!input) throw std::runtime_error("short fixture"); return value;
|
||||||
|
}
|
||||||
|
int main(int argc, char **argv) try {
|
||||||
|
if (argc != 5) return 2;
|
||||||
|
const std::string stage(argv[3]), directory = std::string(argv[2]) + "/";
|
||||||
|
const auto frames = static_cast<size_t>(std::stoul(argv[4]));
|
||||||
|
if (stage != "root" && stage != "body") throw std::runtime_error("stage must be root or body");
|
||||||
|
auto weights = kimodo::detail::ggml_motion_weights::load(argv[1]); if (!weights) throw std::runtime_error(weights.error());
|
||||||
|
const size_t dimension = stage == "root" ? 546 : 545;
|
||||||
|
auto output = kimodo::detail::run_motion_transformer(
|
||||||
|
**weights, stage + "_model.", read(directory + stage + "_input_0.f32"), dimension,
|
||||||
|
read(directory + stage + "_input_2.f32"), read(directory + stage + "_input_4.f32"),
|
||||||
|
read(directory + stage + "_input_5.f32"), 3, frames);
|
||||||
|
if (!output) throw std::runtime_error(output.error());
|
||||||
|
const auto expected = read(directory + stage + "_output.f32");
|
||||||
|
if (output->size() != expected.size()) throw std::runtime_error("shape mismatch");
|
||||||
|
float maximum = 0; double error = 0, reference = 0;
|
||||||
|
for (size_t i = 0; i < expected.size(); ++i) { const float d = (*output)[i] - expected[i]; maximum = std::max(maximum, std::abs(d)); error += double(d)*d; reference += double(expected[i])*expected[i]; }
|
||||||
|
std::printf("%s max_abs=%g rel_l2=%g\n", stage.c_str(), maximum, std::sqrt(error/reference));
|
||||||
|
return maximum < 2.e-3f ? 0 : 1;
|
||||||
|
} catch (const std::exception &error) { std::fprintf(stderr, "%s\n", error.what()); return 1; }
|
||||||
Loading…
Reference in New Issue
Block a user