MCPcopy Create free account
hub / github.com/RL-Align/RL-Kernel / step

Function step

tests/linear_logp_tp.py:378–390  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

376 )
377
378 def step() -> torch.Tensor:
379 out = op(
380 hidden,
381 local_weight,
382 target,
383 local_bias,
384 tp_group=dist.group.WORLD,
385 vocab_start_index=start,
386 global_vocab_size=args.stress_vocab_size,
387 )
388 loss = out.float().mean()
389 loss.backward()
390 return out
391
392 stress_out = None
393

Callers 1

timed_stepFunction · 0.85

Calls 1

backwardMethod · 0.45

Tested by

no test coverage detected